1use super::*;
4use crate::policy::model::{PartialPolicy, TableRuleKey};
5use serde_json::Value as Json;
6use std::fs;
7
8#[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 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
42pub 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 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 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
165pub 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
214pub 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 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 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}