Skip to main content

sequel_mcp/config/
migrate.rs

1//! v1 → v2 migration parsing and mapping, plus the timestamped backup.
2
3use super::*;
4use crate::policy::model::{PartialPolicy, TableRuleKey};
5use serde_json::Value as Json;
6use std::fs;
7
8/// A parsed v1 file (legacy TypeScript schema).
9#[derive(Debug, Clone)]
10pub struct V1Config {
11    pub connections: Vec<V1Connection>,
12    pub default_connection: Option<String>,
13    pub retention: crate::policy::model::RetentionConfig,
14}
15
16#[derive(Debug, Clone)]
17pub struct V1Connection {
18    pub name: String,
19    pub driver: V1Driver,
20    /// Legacy `databasePolicies` as plain `<db>` → partial maps.
21    pub database_policies: std::collections::BTreeMap<String, PartialPolicy>,
22    pub policy: crate::policy::model::Policy,
23    pub mysql: Option<MySqlConnection>,
24    pub sqlite: Option<SqliteConnection>,
25}
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq)]
28pub enum V1Driver {
29    Mysql,
30    Sqlite,
31}
32
33fn str_field(obj: &Json, key: &str) -> Option<String> {
34    obj.get(key).and_then(|v| v.as_str()).map(str::to_string)
35}
36
37fn parse_policy(obj: &Json) -> Result<crate::policy::model::Policy, ConfigError> {
38    serde_json::from_value(obj.clone())
39        .map_err(|e| ConfigError::Malformed(format!("bad policy object: {e}")))
40}
41
42/// Parse a v1 JSON document. Mirrors the legacy zod preprocessing:
43/// missing `driver` defaults to `mysql`; unknown drivers are errors.
44pub fn parse_v1(json: &Json) -> Result<V1Config, ConfigError> {
45    let obj = json
46        .as_object()
47        .ok_or_else(|| ConfigError::Malformed("v1 config must be an object".into()))?;
48    let version = obj.get("version").and_then(|v| v.as_u64()).unwrap_or(1);
49    if version != 1 {
50        return Err(ConfigError::UnsupportedVersion(version));
51    }
52
53    let mut connections = Vec::new();
54    for (i, raw) in obj
55        .get("connections")
56        .map(|v| v.as_array().cloned().unwrap_or_default())
57        .unwrap_or_default()
58        .iter()
59        .enumerate()
60    {
61        let c = raw
62            .as_object()
63            .ok_or_else(|| ConfigError::Malformed(format!("connection #{i} is not an object")))?;
64        let driver = match c.get("driver").and_then(|v| v.as_str()) {
65            None | Some("mysql") => V1Driver::Mysql,
66            Some("sqlite") => V1Driver::Sqlite,
67            Some(other) => {
68                return Err(ConfigError::Malformed(format!(
69                    "connection #{i} has unknown driver {other:?}"
70                )));
71            }
72        };
73        let policy = c
74            .get("policy")
75            .map(parse_policy)
76            .transpose()?
77            .unwrap_or_default();
78        let database_policies = c
79            .get("databasePolicies")
80            .map(parse_partial_map)
81            .transpose()?
82            .unwrap_or_default();
83
84        let (mysql, sqlite) = match driver {
85            V1Driver::Mysql => {
86                let ssh = c.get("ssh").map(parse_ssh).transpose()?;
87                (
88                    Some(MySqlConnection {
89                        name: str_field(raw, "name").unwrap_or_default(),
90                        host: str_field(raw, "host").unwrap_or_default(),
91                        port: raw.get("port").and_then(|v| v.as_u64()).unwrap_or(3306) as u16,
92                        user: str_field(raw, "user").unwrap_or_default(),
93                        database: str_field(raw, "database"),
94                        ssl: raw.get("ssl").and_then(|v| v.as_bool()).unwrap_or(false),
95                        ssl_server_name: str_field(raw, "sslServerName"),
96                        ssl_ca_path: str_field(raw, "sslCaPath"),
97                        ssh,
98                        policy: policy.clone(),
99                        table_policies: TablePolicies::default(),
100                    }),
101                    None,
102                )
103            }
104            V1Driver::Sqlite => (
105                None,
106                Some(SqliteConnection {
107                    name: str_field(raw, "name").unwrap_or_default(),
108                    path: str_field(raw, "path").unwrap_or_default(),
109                    database: str_field(raw, "database").unwrap_or_else(|| "main".into()),
110                    policy: policy.clone(),
111                    table_policies: TablePolicies::default(),
112                }),
113            ),
114        };
115        connections.push(V1Connection {
116            name: str_field(raw, "name").unwrap_or_default(),
117            driver,
118            database_policies,
119            policy,
120            mysql,
121            sqlite,
122        });
123    }
124
125    // Retention with the legacy `auditDays` fan-out.
126    let retention: crate::policy::model::RetentionConfig = obj
127        .get("retention")
128        .map(parse_retention)
129        .transpose()?
130        .unwrap_or_default();
131
132    Ok(V1Config {
133        connections,
134        default_connection: str_field(json, "defaultConnection"),
135        retention,
136    })
137}
138
139fn parse_retention(value: &Json) -> Result<crate::policy::model::RetentionConfig, ConfigError> {
140    let mut v = value.clone();
141    if let Some(obj) = v.as_object_mut() {
142        // Legacy preprocess: `auditDays` fans out to every category when
143        // `retentionDaysByCategory` is absent.
144        if let Some(days) = obj.get("auditDays").and_then(|d| d.as_u64()) {
145            obj.entry("retentionDaysByCategory".to_string())
146                .or_insert_with(|| {
147                    serde_json::json!({
148                        "read": days, "write": days, "ddl": days,
149                        "admin": days, "txCtrl": days
150                    })
151                });
152        }
153        obj.remove("auditDays");
154    }
155    serde_json::from_value(v)
156        .map_err(|e| ConfigError::Malformed(format!("bad retention object: {e}")))
157}
158
159fn parse_ssh(value: &Json) -> Result<SshTunnel, ConfigError> {
160    let ssh: SshTunnel = serde_json::from_value(value.clone())
161        .map_err(|e| ConfigError::Malformed(format!("bad ssh object: {e}")))?;
162    Ok(ssh)
163}
164
165/// Map a parsed v1 config to v2: each `databasePolicies[db]` becomes the
166/// wildcard table rule `db.*`; unset `hostKeyPolicy` becomes `lenient` with
167/// the migrated warning flag.
168pub fn v1_to_v2(v1: V1Config) -> Config {
169    let connections = v1
170        .connections
171        .into_iter()
172        .map(|c| {
173            let mut table_policies = TablePolicies::new();
174            for (db, partial) in &c.database_policies {
175                table_policies.insert(
176                    TableRuleKey::Wildcard {
177                        database: db.clone(),
178                    },
179                    partial.clone(),
180                );
181            }
182            match c.driver {
183                V1Driver::Mysql => {
184                    let mut m = c.mysql.unwrap_or_default();
185                    m.policy = c.policy;
186                    m.table_policies = table_policies;
187                    if let Some(ssh) = &mut m.ssh
188                        && ssh.host_key_policy.is_none()
189                    {
190                        ssh.host_key_policy = Some(SshHostKeyPolicy::Lenient);
191                        ssh.host_key_policy_migrated = true;
192                    }
193                    Connection::Mysql(m)
194                }
195                V1Driver::Sqlite => {
196                    let mut s = c.sqlite.unwrap_or_default();
197                    s.policy = c.policy;
198                    s.table_policies = table_policies;
199                    Connection::Sqlite(s)
200                }
201            }
202        })
203        .collect();
204
205    Config {
206        version: 2,
207        revision: 1,
208        connections,
209        default_connection: v1.default_connection,
210        retention: v1.retention,
211    }
212}
213
214/// Run the on-disk v1 → v2 migration with the documented safety steps.
215pub fn migrate_file(store: &ConfigStore) -> Result<Config, ConfigError> {
216    let _guard = store.lock()?;
217    if !store.path().exists() {
218        return Ok(Config::default());
219    }
220    let raw = fs::read_to_string(store.path())?;
221    let json: Json = serde_json::from_str(&raw)
222        .map_err(|e| ConfigError::Malformed(format!("not valid JSON: {e}")))?;
223    let version = json
224        .get("version")
225        .and_then(|v| v.as_u64())
226        .ok_or_else(|| ConfigError::Malformed("missing numeric version field".into()))?;
227    if version == 2 {
228        return serde_json::from_value(json)
229            .map_err(|e| ConfigError::Malformed(format!("config does not match v2 schema: {e}")));
230    }
231    if version != 1 {
232        return Err(ConfigError::UnsupportedVersion(version));
233    }
234
235    let v1 = parse_v1(&json)?;
236    let v2 = v1_to_v2(v1);
237
238    // Timestamped backup of the original v1 (mode 0600).
239    let stamp = time::OffsetDateTime::now_utc()
240        .format(&time::format_description::well_known::Rfc3339)
241        .unwrap_or_else(|_| "unknown".to_string())
242        .replace(':', "");
243    let backup = store
244        .path()
245        .with_file_name(format!("config.pre-v2.{stamp}.json"));
246    {
247        let mut f = fs::OpenOptions::new()
248            .create(true)
249            .write(true)
250            .truncate(true)
251            .open(&backup)?;
252        f.write_all(raw.as_bytes())?;
253        f.sync_all()?;
254        #[cfg(unix)]
255        {
256            use std::os::unix::fs::PermissionsExt;
257            let _ = fs::set_permissions(&backup, fs::Permissions::from_mode(0o600));
258        }
259    }
260
261    store.persist(&v2)?;
262
263    // Reopen and validate: on failure the v1 original is untouched (the
264    // rename either happened atomically or not at all) and the backup above
265    // remains for manual rollback.
266    let reopened = store.load_locked()?;
267    if reopened.version != 2 {
268        return Err(ConfigError::Validation(
269            "post-migration validation failed".into(),
270        ));
271    }
272    Ok(reopened)
273}
274
275#[cfg(test)]
276mod tests {
277    use super::*;
278
279    #[test]
280    fn legacy_audit_days_fans_out() {
281        let v = serde_json::json!({"auditDays": 14});
282        let r = parse_retention(&v).unwrap();
283        assert_eq!(r.retention_days_by_category.read, 14);
284        assert_eq!(r.retention_days_by_category.admin, 14);
285    }
286
287    #[test]
288    fn driverless_defaults_to_mysql() {
289        let j = serde_json::json!({
290            "version": 1,
291            "connections": [{"name": "c", "host": "h", "user": "u", "policy": {}}]
292        });
293        let v1 = parse_v1(&j).unwrap();
294        assert_eq!(v1.connections[0].driver, V1Driver::Mysql);
295    }
296
297    #[test]
298    fn ssh_policy_unset_becomes_lenient_migrated() {
299        let j = serde_json::json!({
300            "version": 1,
301            "connections": [{
302                "name": "c", "host": "h", "user": "u", "policy": {},
303                "ssh": {"host": "s", "user": "su"}
304            }]
305        });
306        let v2 = v1_to_v2(parse_v1(&j).unwrap());
307        match &v2.connections[0] {
308            Connection::Mysql(m) => {
309                let ssh = m.ssh.as_ref().unwrap();
310                assert_eq!(ssh.host_key_policy, Some(SshHostKeyPolicy::Lenient));
311                assert!(ssh.host_key_policy_migrated);
312            }
313            _ => panic!("expected mysql"),
314        }
315    }
316}