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)]
32#[serde(default)]
33pub struct AppMeta {
34 pub name: String,
37}
38
39impl Default for AppMeta {
40 fn default() -> Self {
41 Self {
42 name: "Laterite".to_string(),
43 }
44 }
45}
46
47#[derive(Debug, Clone, Deserialize)]
50#[serde(default)]
51pub struct BackendConfig {
52 pub secure_cookie: bool,
55 pub timezone: String,
59}
60
61impl Default for BackendConfig {
62 fn default() -> Self {
63 Self {
64 secure_cookie: false,
65 timezone: "UTC".to_string(),
66 }
67 }
68}
69
70fn default_max_connections() -> u32 {
71 10
72}
73
74fn default_acquire_timeout_secs() -> u64 {
75 5
76}
77
78fn default_listen() -> String {
79 "127.0.0.1:8080".to_string()
80}
81
82pub fn load<T: DeserializeOwned>(dir: &Path, env_prefix: &str) -> CoreResult<T> {
90 let app_env = std::env::var("APP_ENV").unwrap_or_else(|_| "development".into());
91 Config::builder()
92 .add_source(File::from(dir.join("default.toml")).required(true))
93 .add_source(File::from(dir.join(format!("{app_env}.toml"))).required(false))
94 .add_source(File::from(dir.join("local.toml")).required(false))
95 .add_source(
96 Environment::with_prefix(env_prefix)
97 .prefix_separator("__")
98 .separator("__"),
99 )
100 .build()
101 .and_then(Config::try_deserialize)
102 .map_err(|e| CoreError::Config(e.to_string()))
103}
104
105#[cfg(test)]
106mod tests {
107 use super::*;
108
109 #[derive(Deserialize)]
110 struct TestConfig {
111 server: ServerConfig,
112 database: DatabaseConfig,
113 }
114
115 #[test]
116 fn loads_defaults_env_overrides_and_serde_defaults() {
117 let dir = tempfile::tempdir().unwrap();
118 std::fs::write(
119 dir.path().join("default.toml"),
120 "[server]\nlisten = \"0.0.0.0:9999\"\n\n[database]\nurl = \"postgres://from-file\"\n",
121 )
122 .unwrap();
123 std::env::set_var("LATERITE_TEST__DATABASE__URL", "postgres://from-env");
124 let cfg: TestConfig = load(dir.path(), "LATERITE_TEST").unwrap();
125 std::env::remove_var("LATERITE_TEST__DATABASE__URL");
126 assert_eq!(cfg.server.listen, "0.0.0.0:9999");
127 assert_eq!(cfg.database.url, "postgres://from-env");
128 assert_eq!(cfg.database.max_connections, 10);
129 assert_eq!(cfg.database.acquire_timeout_secs, 5);
130 }
131
132 #[test]
133 fn backend_config_defaults_and_loads() {
134 #[derive(Deserialize)]
135 struct C {
136 #[serde(default)]
137 backend: BackendConfig,
138 }
139 let dir = tempfile::tempdir().unwrap();
140 std::fs::write(dir.path().join("default.toml"), "").unwrap();
141 let c: C = load(dir.path(), "LATERITE_BE_NONE").unwrap();
142 assert!(!c.backend.secure_cookie);
143 assert_eq!(c.backend.timezone, "UTC");
144
145 std::fs::write(
146 dir.path().join("default.toml"),
147 "[backend]\nsecure_cookie = true\ntimezone = \"Asia/Kolkata\"\n",
148 )
149 .unwrap();
150 let c: C = load(dir.path(), "LATERITE_BE_SET").unwrap();
151 assert!(c.backend.secure_cookie);
152 assert_eq!(c.backend.timezone, "Asia/Kolkata");
153 }
154
155 #[test]
156 fn missing_default_file_is_a_config_error() {
157 let dir = tempfile::tempdir().unwrap();
158 let result: CoreResult<TestConfig> = load(dir.path(), "LATERITE_TEST_MISSING");
159 assert!(matches!(result, Err(CoreError::Config(_))));
160 }
161}