1use std::sync::LazyLock;
11use std::time::Duration;
12
13use crate::utils::server::parse_migrate_startup_timeout;
14use deadpool_postgres::{Manager, ManagerConfig, Pool, RecyclingMethod, Runtime};
15use tokio_postgres::config::Host;
16use tokio_postgres::{Config, NoTls};
17
18fn build_pg_config() -> Result<tokio_postgres::Config, String> {
26 let db_url = std::env::var("DATABASE_URL")
27 .map_err(|_| "DATABASE_URL environment variable not set".to_string())?;
28 let mut pg_cfg = db_url
29 .parse::<tokio_postgres::Config>()
30 .map_err(|e| format!("Invalid DATABASE_URL format: {e}"))?;
31
32 let statement_timeout_secs = std::env::var("STATEMENT_TIMEOUT_SECS")
35 .ok()
36 .and_then(|s| s.parse::<u32>().ok())
37 .unwrap_or(30);
38 pg_cfg.options(format!(
40 "-c statement_timeout={}",
41 statement_timeout_secs * 1000
42 ));
43
44 Ok(pg_cfg)
45}
46
47pub fn validate_database_url() -> Result<(), String> {
55 build_pg_config()?;
56
57 if let Ok(s) = std::env::var("DB_POOL_SIZE") {
59 match s.parse::<usize>() {
60 Ok(n) if n > 0 => {}
61 Ok(_) => return Err("DB_POOL_SIZE is not a positive integer".to_string()),
62 Err(e) => return Err(format!("Invalid DB_POOL_SIZE value: {e}")),
63 }
64 }
65 Ok(())
66}
67
68pub static DB_POOL: LazyLock<Pool> = LazyLock::new(|| {
76 let pg_cfg = build_pg_config()
78 .expect("DATABASE_URL should have been validated at startup; validate_database_url() was not called");
79
80 let mgr_cfg = ManagerConfig {
84 recycling_method: RecyclingMethod::Fast,
85 };
86 let mgr = Manager::from_config(pg_cfg, NoTls, mgr_cfg);
87
88 Pool::builder(mgr)
89 .max_size(
90 std::env::var("DB_POOL_SIZE")
91 .ok()
92 .and_then(|s| s.parse().ok())
93 .unwrap_or(20),
94 )
95 .wait_timeout(Some(Duration::from_secs(10)))
96 .create_timeout(Some(Duration::from_secs(10)))
97 .recycle_timeout(Some(Duration::from_secs(5)))
98 .runtime(Runtime::Tokio1)
99 .build()
100 .expect("Failed to create database connection pool")
101});
102
103pub async fn get_conn() -> Result<deadpool_postgres::Object, deadpool_postgres::PoolError> {
113 use rand::TryRng;
114
115 let mut last_err = None;
116 for attempt in 0..=crate::db::retry::MAX_RETRIES {
117 match DB_POOL.get().await {
118 Ok(conn) => return Ok(conn),
119 Err(e) => {
120 let is_timeout = matches!(e, deadpool_postgres::PoolError::Timeout(_));
123 last_err = Some(e);
124 if !is_timeout && attempt < crate::db::retry::MAX_RETRIES {
125 let Ok(random_bits) = rand::rngs::SysRng.try_next_u64() else {
126 break;
128 };
129 let jitter = (random_bits >> 11) as f64 * (1.0 / (1u64 << 53) as f64);
131 let delay = crate::db::retry::backoff_for(attempt, jitter);
132 tracing::warn!(
133 "DB connection attempt {} failed (backend error), retrying in {:?}: {:?}",
134 attempt + 1,
135 delay,
136 last_err.as_ref().unwrap(),
137 );
138 tokio::time::sleep(delay).await;
139 } else if is_timeout {
140 break;
142 }
143 }
144 }
145 }
146 Err(last_err.unwrap())
147}
148
149pub async fn get_conn_for_startup(
162) -> Result<deadpool_postgres::Object, deadpool_postgres::PoolError> {
163 let timeout_secs = parse_migrate_startup_timeout();
164
165 let deadline = tokio::time::Instant::now() + Duration::from_secs(timeout_secs);
166 let retry_interval = Duration::from_millis(500);
167
168 let mut attempt = 0u32;
169 loop {
170 attempt += 1;
171 match DB_POOL.get().await {
172 Ok(conn) => {
173 if attempt > 1 {
174 tracing::info!("connected to database after {} attempt(s)", attempt);
175 }
176 return Ok(conn);
177 }
178 Err(e) => {
179 let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
180 if remaining.is_zero() {
181 return Err(e);
182 }
183 tracing::warn!(
184 "startup DB connection attempt {} failed, ~{}s remaining until giving up: {:?}",
185 attempt,
186 remaining.as_secs(),
187 e,
188 );
189 let sleep = std::cmp::min(retry_interval, remaining);
191 tokio::time::sleep(sleep).await;
192 }
193 }
194 }
195}
196
197pub async fn ensure_database() -> Result<(), String> {
217 let pg_cfg = build_pg_config()?;
219 let db_name = pg_cfg
220 .get_dbname()
221 .or_else(|| pg_cfg.get_user())
222 .map(|s| s.to_string());
223
224 let db_name = match db_name {
225 Some(name) => name,
226 None => {
227 tracing::warn!(
228 "could not determine target database name from DATABASE_URL; \
229 skipping auto-create (letting normal connect path surface any error)"
230 );
231 return Ok(());
232 }
233 };
234
235 if !is_simple_ident(&db_name) {
238 tracing::warn!(
239 "target database name {:?} is not a simple identifier; \
240 skipping auto-create (letting normal connect path surface any error)",
241 db_name
242 );
243 return Ok(());
244 }
245
246 let timeout_secs = parse_migrate_startup_timeout();
248 let deadline = tokio::time::Instant::now() + Duration::from_secs(timeout_secs);
249 let retry_interval = Duration::from_millis(500);
250
251 let (client, connection) = loop {
252 let admin_cfg = build_admin_config()?;
253 match admin_cfg.connect(NoTls).await {
254 Ok(joined) => break joined,
255 Err(e) => {
256 let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
257 if remaining.is_zero() {
258 return Err(format!(
259 "could not connect to 'postgres' maintenance database within {timeout_secs}s: {}",
260 crate::db::format_with_sources(&e)
261 ));
262 }
263 tracing::warn!(
264 "ensure_database: connect to 'postgres' failed, ~{}s remaining: {}",
265 remaining.as_secs(),
266 crate::db::format_with_sources(&e)
267 );
268 tokio::time::sleep(std::cmp::min(retry_interval, remaining)).await;
269 }
270 }
271 };
272 tokio::spawn(async move {
274 if let Err(e) = connection.await {
275 tracing::warn!(
276 "postgres maintenance connection ended: {}",
277 crate::db::format_with_sources(&e)
278 );
279 }
280 });
281
282 let exists: bool = client
284 .query_one(
285 "SELECT EXISTS(SELECT 1 FROM pg_database WHERE datname = $1)",
286 &[&db_name],
287 )
288 .await
289 .map_err(|e| {
290 format!(
291 "failed to query pg_database: {}",
292 crate::db::format_with_sources(&e)
293 )
294 })?
295 .get(0);
296
297 if exists {
298 tracing::info!("target database {:?} already exists", db_name);
299 return Ok(());
300 }
301
302 tracing::info!("target database {:?} does not exist, creating", db_name);
304 let stmt = format!("CREATE DATABASE {db_name}");
305 client.batch_execute(&stmt).await.map_err(|e| {
306 format!(
307 "failed to create database {db_name:?}: {}",
308 crate::db::format_with_sources(&e)
309 )
310 })?;
311 tracing::info!("created database {:?}", db_name);
312 Ok(())
313}
314
315fn build_admin_config() -> Result<Config, String> {
321 let src = build_pg_config()?;
322 let mut dst = Config::new();
323
324 if let Some(user) = src.get_user() {
325 dst.user(user);
326 }
327 if let Some(password) = src.get_password() {
328 dst.password(password);
329 }
330 let hosts = src.get_hosts();
331 let ports = src.get_ports();
332 for (i, host) in hosts.iter().enumerate() {
333 let port = ports.get(i).copied().unwrap_or(5432);
335 match host {
336 Host::Tcp(h) => {
337 dst.host(h);
338 dst.port(port);
339 }
340 Host::Unix(p) => {
341 dst.host_path(p);
342 }
343 };
344 }
345 dst.dbname("postgres");
347 Ok(dst)
348}
349
350fn is_simple_ident(s: &str) -> bool {
355 let mut chars = s.chars();
356 match chars.next() {
357 Some(c) if c.is_ascii_alphabetic() || c == '_' => {}
358 _ => return false,
359 }
360 chars.all(|c| c.is_ascii_alphanumeric() || c == '_')
361}
362
363#[cfg(test)]
364mod tests {
365 use super::is_simple_ident;
366
367 #[test]
368 fn simple_ident_accepts_valid_names() {
369 assert!(is_simple_ident("yggdrasil"));
370 assert!(is_simple_ident("_ygg"));
371 assert!(is_simple_ident("db_1"));
372 assert!(is_simple_ident("YggDrasil09"));
373 }
374
375 #[test]
376 fn simple_ident_rejects_invalid_names() {
377 assert!(!is_simple_ident(""));
378 assert!(!is_simple_ident("my-db")); assert!(!is_simple_ident("9db")); assert!(!is_simple_ident("db name")); assert!(!is_simple_ident("db\"; --")); assert!(!is_simple_ident("db.name")); }
384}