1use std::path::Path;
4
5use config::{Config, Environment, File};
6use serde::de::DeserializeOwned;
7use serde::Deserialize;
8
9use crate::error::{CoreError, CoreResult};
10
11#[derive(Debug, Clone, Deserialize)]
13pub struct DatabaseConfig {
14 pub url: String,
15 #[serde(default = "default_max_connections")]
16 pub max_connections: u32,
17 #[serde(default = "default_acquire_timeout_secs")]
18 pub acquire_timeout_secs: u64,
19}
20
21#[derive(Debug, Clone, Deserialize)]
23pub struct ServerConfig {
24 #[serde(default = "default_listen")]
25 pub listen: String,
26}
27
28#[derive(Debug, Clone, Deserialize)]
31#[serde(default)]
32pub struct BackendConfig {
33 pub secure_cookie: bool,
36 pub timezone: String,
40}
41
42impl Default for BackendConfig {
43 fn default() -> Self {
44 Self {
45 secure_cookie: false,
46 timezone: "UTC".to_string(),
47 }
48 }
49}
50
51fn default_max_connections() -> u32 {
52 10
53}
54
55fn default_acquire_timeout_secs() -> u64 {
56 5
57}
58
59fn default_listen() -> String {
60 "127.0.0.1:8080".to_string()
61}
62
63pub fn load<T: DeserializeOwned>(dir: &Path, env_prefix: &str) -> CoreResult<T> {
71 let app_env = std::env::var("APP_ENV").unwrap_or_else(|_| "development".into());
72 Config::builder()
73 .add_source(File::from(dir.join("default.toml")).required(true))
74 .add_source(File::from(dir.join(format!("{app_env}.toml"))).required(false))
75 .add_source(File::from(dir.join("local.toml")).required(false))
76 .add_source(
77 Environment::with_prefix(env_prefix)
78 .prefix_separator("__")
79 .separator("__"),
80 )
81 .build()
82 .and_then(Config::try_deserialize)
83 .map_err(|e| CoreError::Config(e.to_string()))
84}
85
86#[cfg(test)]
87mod tests {
88 use super::*;
89
90 #[derive(Deserialize)]
91 struct TestConfig {
92 server: ServerConfig,
93 database: DatabaseConfig,
94 }
95
96 #[test]
97 fn loads_defaults_env_overrides_and_serde_defaults() {
98 let dir = tempfile::tempdir().unwrap();
99 std::fs::write(
100 dir.path().join("default.toml"),
101 "[server]\nlisten = \"0.0.0.0:9999\"\n\n[database]\nurl = \"postgres://from-file\"\n",
102 )
103 .unwrap();
104 std::env::set_var("LATERITE_TEST__DATABASE__URL", "postgres://from-env");
105 let cfg: TestConfig = load(dir.path(), "LATERITE_TEST").unwrap();
106 std::env::remove_var("LATERITE_TEST__DATABASE__URL");
107 assert_eq!(cfg.server.listen, "0.0.0.0:9999");
108 assert_eq!(cfg.database.url, "postgres://from-env");
109 assert_eq!(cfg.database.max_connections, 10);
110 assert_eq!(cfg.database.acquire_timeout_secs, 5);
111 }
112
113 #[test]
114 fn backend_config_defaults_and_loads() {
115 #[derive(Deserialize)]
116 struct C {
117 #[serde(default)]
118 backend: BackendConfig,
119 }
120 let dir = tempfile::tempdir().unwrap();
121 std::fs::write(dir.path().join("default.toml"), "").unwrap();
122 let c: C = load(dir.path(), "LATERITE_BE_NONE").unwrap();
123 assert!(!c.backend.secure_cookie);
124 assert_eq!(c.backend.timezone, "UTC");
125
126 std::fs::write(
127 dir.path().join("default.toml"),
128 "[backend]\nsecure_cookie = true\ntimezone = \"Asia/Kolkata\"\n",
129 )
130 .unwrap();
131 let c: C = load(dir.path(), "LATERITE_BE_SET").unwrap();
132 assert!(c.backend.secure_cookie);
133 assert_eq!(c.backend.timezone, "Asia/Kolkata");
134 }
135
136 #[test]
137 fn missing_default_file_is_a_config_error() {
138 let dir = tempfile::tempdir().unwrap();
139 let result: CoreResult<TestConfig> = load(dir.path(), "LATERITE_TEST_MISSING");
140 assert!(matches!(result, Err(CoreError::Config(_))));
141 }
142}