1pub 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 #[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 pub ssl_ca_path: Option<String>,
148 pub ssh: Option<SshTunnel>,
149 pub policy: Policy,
150 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#[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 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 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#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
329#[serde(rename_all = "camelCase", default)]
330pub struct Config {
331 pub version: u32,
332 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#[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 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 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 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 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
565pub(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
586pub 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 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}