Skip to main content

sequel_mcp/config/
mod.rs

1//! Configuration: connection model, v2 store with locked atomic writes,
2//! and the v1 → v2 migration.
3
4pub mod migrate;
5
6use crate::policy::model::{PartialPolicy, Policy, RetentionConfig, TablePolicies, TableRuleKey};
7use serde::{Deserialize, Serialize};
8use serde_json::Value as Json;
9use std::collections::BTreeMap;
10use std::fs;
11use std::io::Write;
12use std::path::{Path, PathBuf};
13use std::sync::Arc;
14use thiserror::Error;
15
16use crate::app::paths;
17
18#[derive(Debug, Error)]
19pub enum ConfigError {
20    #[error("config file is malformed: {0}")]
21    Malformed(String),
22    #[error("unsupported config version: {0}")]
23    UnsupportedVersion(u64),
24    #[error("io error: {0}")]
25    Io(#[from] std::io::Error),
26    #[error("serialization error: {0}")]
27    Serde(#[from] serde_json::Error),
28    #[error(
29        "revision conflict: config changed on disk (expected revision {expected}, found {found})"
30    )]
31    RevisionConflict { expected: u64, found: u64 },
32    #[error("connection name {0:?} is invalid (allowed: A-Za-z0-9 _-:. up to 128 chars)")]
33    InvalidConnectionName(String),
34    #[error("validation error: {0}")]
35    Validation(String),
36}
37
38#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
39#[serde(rename_all = "lowercase")]
40pub enum SshHostKeyPolicy {
41    Lenient,
42    Strict,
43}
44
45#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
46#[serde(rename_all = "lowercase")]
47pub enum BridgeTool {
48    Nc,
49    Ncat,
50    Socat,
51}
52
53impl BridgeTool {
54    pub fn as_str(&self) -> &'static str {
55        match self {
56            BridgeTool::Nc => "nc",
57            BridgeTool::Ncat => "ncat",
58            BridgeTool::Socat => "socat",
59        }
60    }
61
62    pub fn parse(s: &str) -> Option<Self> {
63        match s {
64            "nc" => Some(BridgeTool::Nc),
65            "ncat" => Some(BridgeTool::Ncat),
66            "socat" => Some(BridgeTool::Socat),
67            _ => None,
68        }
69    }
70}
71
72#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
73#[serde(rename_all = "camelCase", default)]
74pub struct SshDocker {
75    pub container: String,
76    pub bridge_tool: BridgeTool,
77}
78
79impl Default for SshDocker {
80    fn default() -> Self {
81        Self {
82            container: String::new(),
83            bridge_tool: BridgeTool::Nc,
84        }
85    }
86}
87
88impl SshDocker {
89    pub fn validate(&self) -> Result<(), ConfigError> {
90        crate::sql::docker::validate_container_name(&self.container)
91            .map_err(ConfigError::Validation)
92    }
93}
94
95#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
96#[serde(rename_all = "camelCase", default)]
97pub struct SshTunnel {
98    pub host: String,
99    pub port: u16,
100    pub user: String,
101    pub auth_method: SshAuthMethod,
102    pub private_key_path: Option<String>,
103    pub docker: Option<SshDocker>,
104    pub host_key_policy: Option<SshHostKeyPolicy>,
105    /// Set when `host_key_policy` was inherited unset from v1 (lenient with
106    /// a prominent warning rather than silently trusted).
107    #[serde(skip_serializing_if = "std::ops::Not::not")]
108    pub host_key_policy_migrated: bool,
109    pub known_hosts_path: Option<String>,
110}
111
112impl Default for SshTunnel {
113    fn default() -> Self {
114        Self {
115            host: String::new(),
116            port: 22,
117            user: String::new(),
118            auth_method: SshAuthMethod::Key,
119            private_key_path: None,
120            docker: None,
121            host_key_policy: None,
122            host_key_policy_migrated: false,
123            known_hosts_path: None,
124        }
125    }
126}
127
128#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
129#[serde(rename_all = "lowercase")]
130pub enum SshAuthMethod {
131    Password,
132    Key,
133}
134
135#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
136#[serde(rename_all = "camelCase", default)]
137pub struct MySqlConnection {
138    pub name: String,
139    pub host: String,
140    pub port: u16,
141    pub user: String,
142    pub database: Option<String>,
143    pub ssl: bool,
144    pub ssl_server_name: Option<String>,
145    /// Optional PEM/DER file of a private CA the server certificate is
146    /// verified against (merged with the system roots).
147    pub ssl_ca_path: Option<String>,
148    pub ssh: Option<SshTunnel>,
149    pub policy: Policy,
150    /// v2 layer-2 rules (exact or wildcard). Migrated v1 `databasePolicies`
151    /// appear here as `db.*` wildcards.
152    pub table_policies: TablePolicies,
153}
154
155impl Default for MySqlConnection {
156    fn default() -> Self {
157        Self {
158            name: String::new(),
159            host: String::new(),
160            port: 3306,
161            user: String::new(),
162            database: None,
163            ssl: false,
164            ssl_server_name: None,
165            ssl_ca_path: None,
166            ssh: None,
167            policy: Policy::default(),
168            table_policies: TablePolicies::default(),
169        }
170    }
171}
172
173#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
174#[serde(rename_all = "camelCase", default)]
175pub struct SqliteConnection {
176    pub name: String,
177    pub path: String,
178    pub database: String,
179    pub policy: Policy,
180    pub table_policies: TablePolicies,
181}
182
183impl Default for SqliteConnection {
184    fn default() -> Self {
185        Self {
186            name: String::new(),
187            path: String::new(),
188            database: "main".to_string(),
189            policy: Policy::default(),
190            table_policies: TablePolicies::default(),
191        }
192    }
193}
194
195#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
196#[serde(tag = "driver", rename_all = "lowercase")]
197// Boxing the MySQL variant would ripple through every match site for a
198// type that is constructed rarely and cloned per request anyway.
199#[allow(clippy::large_enum_variant)]
200pub enum Connection {
201    Mysql(MySqlConnection),
202    Sqlite(SqliteConnection),
203}
204
205impl Connection {
206    pub fn name(&self) -> &str {
207        match self {
208            Connection::Mysql(c) => &c.name,
209            Connection::Sqlite(c) => &c.name,
210        }
211    }
212
213    pub fn database(&self) -> Option<&str> {
214        match self {
215            Connection::Mysql(c) => c.database.as_deref(),
216            Connection::Sqlite(c) => Some(&c.database),
217        }
218    }
219
220    pub fn policy(&self) -> &Policy {
221        match self {
222            Connection::Mysql(c) => &c.policy,
223            Connection::Sqlite(c) => &c.policy,
224        }
225    }
226
227    pub fn table_policies(&self) -> &TablePolicies {
228        match self {
229            Connection::Mysql(c) => &c.table_policies,
230            Connection::Sqlite(c) => &c.table_policies,
231        }
232    }
233
234    pub fn is_mysql(&self) -> bool {
235        matches!(self, Connection::Mysql(_))
236    }
237
238    pub fn validate(&self) -> Result<(), ConfigError> {
239        let name = self.name();
240        let valid = !name.is_empty()
241            && name.len() <= 128
242            && name
243                .chars()
244                .all(|c| c.is_ascii_alphanumeric() || matches!(c, ' ' | '_' | '-' | ':' | '.'));
245        if !valid {
246            return Err(ConfigError::InvalidConnectionName(name.to_string()));
247        }
248        self.policy()
249            .validate()
250            .map_err(|e| ConfigError::Validation(e.into()))?;
251        // Table rules must respect the same bounds as the baseline
252        // (review finding: rules previously bypassed validation).
253        for rule in self.table_policies().values() {
254            rule.validate()
255                .map_err(|e| ConfigError::Validation(e.into()))?;
256        }
257        if let Connection::Mysql(c) = self {
258            if c.host.is_empty() {
259                return Err(ConfigError::Validation(
260                    "mysql host must not be empty".into(),
261                ));
262            }
263            if c.user.is_empty() {
264                return Err(ConfigError::Validation(
265                    "mysql user must not be empty".into(),
266                ));
267            }
268            if let Some(name) = &c.ssl_server_name
269                && (name.is_empty() || name.len() > 253)
270            {
271                return Err(ConfigError::Validation(
272                    "sslServerName must be 1..=253 chars".into(),
273                ));
274            }
275            if let Some(ca) = &c.ssl_ca_path
276                && (ca.is_empty() || ca.len() > 4096)
277            {
278                return Err(ConfigError::Validation(
279                    "sslCaPath must be 1..=4096 chars".into(),
280                ));
281            }
282            if let Some(ssh) = &c.ssh {
283                if ssh.host.is_empty() || ssh.user.is_empty() {
284                    return Err(ConfigError::Validation(
285                        "ssh host and user must not be empty".into(),
286                    ));
287                }
288                if ssh.auth_method == SshAuthMethod::Key && ssh.private_key_path.is_none() {
289                    return Err(ConfigError::Validation(
290                        "ssh key auth requires privateKeyPath".into(),
291                    ));
292                }
293                // Review finding: a relative knownHostsPath would resolve
294                // against the process CWD (often client-controlled),
295                // letting an attacker-placed known_hosts satisfy even
296                // strict mode. Require absolute paths (a leading `~` is
297                // fine — the loader expands it).
298                if let Some(p) = &ssh.known_hosts_path
299                    && !std::path::Path::new(p).is_absolute()
300                    && !p.starts_with('~')
301                {
302                    return Err(ConfigError::Validation(format!(
303                        "ssh knownHostsPath must be absolute (got {p:?})"
304                    )));
305                }
306                if let Some(d) = &ssh.docker {
307                    d.validate()?;
308                }
309            }
310        }
311        if let Connection::Sqlite(c) = self {
312            if c.path.is_empty() || c.path.len() > 4096 {
313                return Err(ConfigError::Validation(
314                    "sqlite path must be 1..=4096 chars".into(),
315                ));
316            }
317            if c.database.is_empty() || c.database.len() > 64 {
318                return Err(ConfigError::Validation(
319                    "sqlite database (schema) must be 1..=64 chars".into(),
320                ));
321            }
322        }
323        Ok(())
324    }
325}
326
327/// v2 top-level configuration.
328#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
329#[serde(rename_all = "camelCase", default)]
330pub struct Config {
331    pub version: u32,
332    /// Compare-and-swap revision, bumped on every persisted change.
333    pub revision: u64,
334    pub connections: Vec<Connection>,
335    pub default_connection: Option<String>,
336    pub retention: RetentionConfig,
337}
338
339impl Default for Config {
340    fn default() -> Self {
341        Self {
342            version: 2,
343            revision: 0,
344            connections: Vec::new(),
345            default_connection: None,
346            retention: RetentionConfig::default(),
347        }
348    }
349}
350
351impl Config {
352    pub fn get(&self, name: &str) -> Option<&Connection> {
353        self.connections.iter().find(|c| c.name() == name)
354    }
355
356    pub fn resolve(&self, explicit: Option<&str>) -> Option<&Connection> {
357        match explicit {
358            Some(name) => self.get(name),
359            None => self
360                .default_connection
361                .as_deref()
362                .and_then(|def| self.get(def)),
363        }
364    }
365}
366
367/// Handle providing read/write access to the on-disk config with
368/// inter-process locking and revision-checked writes.
369#[derive(Debug, Clone)]
370pub struct ConfigStore {
371    path: Arc<PathBuf>,
372}
373
374impl ConfigStore {
375    pub fn new() -> Self {
376        Self {
377            path: Arc::new(paths::config_path()),
378        }
379    }
380
381    pub fn with_path(path: PathBuf) -> Self {
382        Self {
383            path: Arc::new(path),
384        }
385    }
386
387    pub fn path(&self) -> &Path {
388        &self.path
389    }
390
391    /// Load the config, transparently migrating a v1 file if found.
392    pub fn load(&self) -> Result<Config, ConfigError> {
393        if !self.path.exists() {
394            return Ok(Config::default());
395        }
396        let raw = fs::read_to_string(&*self.path)?;
397        let json: Json = serde_json::from_str(&raw)
398            .map_err(|e| ConfigError::Malformed(format!("not valid JSON: {e}")))?;
399        let version = json
400            .get("version")
401            .and_then(|v| v.as_u64())
402            .ok_or_else(|| ConfigError::Malformed("missing numeric version field".into()))?;
403        match version {
404            1 => {
405                let v1 = migrate::parse_v1(&json)?;
406                let v2 = migrate::v1_to_v2(v1);
407                Ok(v2)
408            }
409            2 => serde_json::from_value(json).map_err(|e| {
410                ConfigError::Malformed(format!("config does not match v2 schema: {e}"))
411            }),
412            other => Err(ConfigError::UnsupportedVersion(other)),
413        }
414    }
415
416    /// Read-modify-write under the inter-process lock; the closure receives
417    /// the current config and returns the next one. Revision is bumped and
418    /// checked, so concurrent writers cannot silently overwrite each other.
419    pub fn update<T>(
420        &self,
421        expected_revision: u64,
422        mutate: impl FnOnce(&mut Config) -> Result<T, ConfigError>,
423    ) -> Result<T, ConfigError> {
424        let _guard = self.lock()?;
425        let mut cfg = self.load_locked()?;
426        if cfg.version == 1 {
427            // v1 on disk: take the timestamped backup HERE (review
428            // finding: `migrate_file` — the only code that created the
429            // backup — is never called, so this path used to persist v2
430            // over v1 with no rollback artifact), then migrate
431            // in-memory before mutating.
432            let raw = fs::read_to_string(self.path())?;
433            let stamp = time::OffsetDateTime::now_utc()
434                .format(&time::format_description::well_known::Rfc3339)
435                .unwrap_or_else(|_| "unknown".to_string())
436                .replace(':', "");
437            let backup = self
438                .path()
439                .with_file_name(format!("config.pre-v2.{stamp}.json"));
440            {
441                use std::io::Write;
442                let mut f = fs::OpenOptions::new()
443                    .create(true)
444                    .write(true)
445                    .truncate(true)
446                    .open(&backup)?;
447                f.write_all(raw.as_bytes())?;
448                f.sync_all()?;
449                #[cfg(unix)]
450                {
451                    use std::os::unix::fs::PermissionsExt;
452                    let _ = fs::set_permissions(&backup, fs::Permissions::from_mode(0o600));
453                }
454            }
455            cfg = migrate::v1_to_v2(migrate::parse_v1(&serde_json::to_value(&cfg)?)?);
456        }
457        if cfg.revision != expected_revision {
458            return Err(ConfigError::RevisionConflict {
459                expected: expected_revision,
460                found: cfg.revision,
461            });
462        }
463        let out = mutate(&mut cfg)?;
464        cfg.revision += 1;
465        for c in &cfg.connections {
466            c.validate()?;
467        }
468        if let Some(def) = &cfg.default_connection
469            && !cfg.connections.iter().any(|c| c.name() == def)
470        {
471            return Err(ConfigError::Validation(format!(
472                "default connection {def:?} does not exist"
473            )));
474        }
475        self.persist(&cfg)?;
476        Ok(out)
477    }
478
479    fn lock(&self) -> Result<fs::File, ConfigError> {
480        if let Some(parent) = self.path.parent() {
481            fs::create_dir_all(parent).map_err(|e| ConfigError::Io(std::io::Error::other(e)))?;
482            set_dir_mode_0700(parent);
483        }
484        let lock_path = self.path.with_extension("lock");
485        let file = fs::OpenOptions::new()
486            .create(true)
487            .truncate(false)
488            .write(true)
489            .open(&lock_path)?;
490        file.lock().map_err(std::io::Error::other)?;
491        Ok(file)
492    }
493
494    fn load_locked(&self) -> Result<Config, ConfigError> {
495        if !self.path.exists() {
496            return Ok(Config::default());
497        }
498        let raw = fs::read_to_string(&*self.path)?;
499        let json: Json = serde_json::from_str(&raw)
500            .map_err(|e| ConfigError::Malformed(format!("not valid JSON: {e}")))?;
501        let version = json
502            .get("version")
503            .and_then(|v| v.as_u64())
504            .ok_or_else(|| ConfigError::Malformed("missing numeric version field".into()))?;
505        match version {
506            1 => Ok(migrate::v1_to_v2(migrate::parse_v1(&json)?)),
507            2 => serde_json::from_value(json).map_err(|e| {
508                ConfigError::Malformed(format!("config does not match v2 schema: {e}"))
509            }),
510            other => Err(ConfigError::UnsupportedVersion(other)),
511        }
512    }
513
514    /// Persist atomically: temp file (0600) → fsync → rename → fsync parent.
515    pub(crate) fn persist(&self, cfg: &Config) -> Result<(), ConfigError> {
516        let parent = self
517            .path
518            .parent()
519            .ok_or_else(|| ConfigError::Validation("config path has no parent".into()))?;
520        fs::create_dir_all(parent).map_err(std::io::Error::other)?;
521        set_dir_mode_0700(parent);
522        let mut json = serde_json::to_string_pretty(cfg)?;
523        json.push('\n');
524        let tmp = self
525            .path
526            .with_extension(format!("tmp.{}", std::process::id()));
527        {
528            let mut f = fs::OpenOptions::new()
529                .create(true)
530                .write(true)
531                .truncate(true)
532                .open(&tmp)?;
533            #[cfg(unix)]
534            {
535                use std::os::unix::fs::PermissionsExt;
536                fs::set_permissions(&tmp, fs::Permissions::from_mode(0o600))
537                    .map_err(std::io::Error::other)?;
538            }
539            f.write_all(json.as_bytes())?;
540            f.sync_all()?;
541        }
542        fs::rename(&tmp, &*self.path)?;
543        if let Ok(dir) = fs::File::open(parent) {
544            let _ = dir.sync_all();
545        }
546        Ok(())
547    }
548}
549
550impl Default for ConfigStore {
551    fn default() -> Self {
552        Self::new()
553    }
554}
555
556#[cfg(unix)]
557fn set_dir_mode_0700(dir: &Path) {
558    use std::os::unix::fs::PermissionsExt;
559    let _ = fs::set_permissions(dir, fs::Permissions::from_mode(0o700));
560}
561
562#[cfg(not(unix))]
563fn set_dir_mode_0700(_dir: &Path) {}
564
565/// Partial policy map parsed from legacy `databasePolicies` JSON objects.
566pub(crate) fn parse_partial_map(
567    value: &Json,
568) -> Result<BTreeMap<String, PartialPolicy>, ConfigError> {
569    let obj = value
570        .as_object()
571        .ok_or_else(|| ConfigError::Malformed("databasePolicies must be an object".into()))?;
572    let mut out = BTreeMap::new();
573    for (k, v) in obj {
574        if k.is_empty() || k.len() > 64 {
575            return Err(ConfigError::Malformed(format!(
576                "database policy key {k:?} out of bounds"
577            )));
578        }
579        let partial: PartialPolicy = serde_json::from_value(v.clone())
580            .map_err(|e| ConfigError::Malformed(format!("bad partial policy for {k:?}: {e}")))?;
581        out.insert(k.clone(), partial);
582    }
583    Ok(out)
584}
585
586/// Helper for the wrapper tools: read a `db.*` rule as the legacy shape.
587pub fn wildcard_rule<'a>(connection: &'a Connection, database: &str) -> Option<&'a PartialPolicy> {
588    connection.table_policies().get(&TableRuleKey::Wildcard {
589        database: database.to_string(),
590    })
591}
592
593#[cfg(test)]
594mod tests {
595    use super::*;
596    use crate::policy::model::{PolicyAction, TableId};
597
598    fn tmp_store() -> (tempfile::TempDir, ConfigStore) {
599        let dir = tempfile::tempdir().unwrap();
600        let store = ConfigStore::with_path(dir.path().join("config.json"));
601        (dir, store)
602    }
603
604    #[test]
605    fn pristine_install_loads_default() {
606        let (_dir, store) = tmp_store();
607        let cfg = store.load().unwrap();
608        assert_eq!(cfg.version, 2);
609        assert!(cfg.connections.is_empty());
610        assert_eq!(cfg.revision, 0);
611    }
612
613    #[test]
614    fn v1_config_migrates_to_wildcard_rules() {
615        let (dir, store) = tmp_store();
616        let v1 = serde_json::json!({
617            "version": 1,
618            "defaultConnection": "c1",
619            "connections": [{
620                "name": "c1",
621                "host": "db.example.invalid",
622                "user": "u",
623                "policy": {
624                    "read": "allow", "write": "deny", "ddl": "deny",
625                    "admin": "deny", "txCtrl": "allow",
626                    "rowCap": 1000, "stmtTimeoutMs": 10000,
627                    "requireTouchID": false, "maxBackupRows": 10000,
628                    "maxBackupBytes": 52428800, "onBackupOverflow": "abort"
629                },
630                "databasePolicies": { "app": { "write": "confirm" } }
631            }],
632            "retention": {}
633        });
634        fs::write(store.path(), serde_json::to_string(&v1).unwrap()).unwrap();
635        let cfg = store.load().unwrap();
636        assert_eq!(cfg.version, 2);
637        let conn = cfg.get("c1").unwrap();
638        let rule = conn
639            .table_policies()
640            .get(&TableRuleKey::Wildcard {
641                database: "app".into(),
642            })
643            .unwrap();
644        assert_eq!(rule.write, Some(PolicyAction::Confirm));
645        // driver-less v1 connection becomes MySQL
646        assert!(conn.is_mysql());
647        let _ = dir;
648    }
649
650    #[test]
651    fn update_bumps_revision_and_conflicts() {
652        let (_dir, store) = tmp_store();
653        let cfg = store.load().unwrap();
654        let () = store
655            .update(cfg.revision, |c| {
656                c.connections.push(Connection::Sqlite(SqliteConnection {
657                    name: "s1".into(),
658                    path: "/tmp/app.sqlite".into(),
659                    ..SqliteConnection::default()
660                }));
661                Ok(())
662            })
663            .unwrap();
664        let cfg2 = store.load().unwrap();
665        assert_eq!(cfg2.revision, 1);
666        let err = store.update(0, |_| Ok(())).unwrap_err();
667        assert!(matches!(err, ConfigError::RevisionConflict { .. }));
668    }
669
670    #[test]
671    fn malformed_json_is_an_error() {
672        let (dir, store) = tmp_store();
673        fs::write(store.path(), "{ not json").unwrap();
674        assert!(matches!(
675            store.load().unwrap_err(),
676            ConfigError::Malformed(_)
677        ));
678        let _ = dir;
679    }
680
681    #[test]
682    fn unknown_version_rejected() {
683        let (dir, store) = tmp_store();
684        fs::write(store.path(), r#"{"version": 99}"#).unwrap();
685        assert!(matches!(
686            store.load().unwrap_err(),
687            ConfigError::UnsupportedVersion(99)
688        ));
689        let _ = dir;
690    }
691
692    #[test]
693    fn exact_rule_lookup_and_validate() {
694        let mut conn = MySqlConnection {
695            name: "c1".into(),
696            host: "h".into(),
697            user: "u".into(),
698            ..MySqlConnection::default()
699        };
700        conn.table_policies.insert(
701            TableRuleKey::Exact(TableId::new("app", "jobs")),
702            PartialPolicy {
703                write: Some(PolicyAction::Allow),
704                ..PartialPolicy::default()
705            },
706        );
707        let c = Connection::Mysql(conn);
708        c.validate().unwrap();
709    }
710}