use windmill_common::{ error::{self, Error}, get_database_url, DatabaseUrl, }; // Single source of truth in windmill_common so the DB-health sizing guidance // (windmill-api/src/db_health.rs) and the actual pool sizing here can't drift. pub use windmill_common::{ DEFAULT_MAX_CONNECTIONS_INDEXER, DEFAULT_MAX_CONNECTIONS_SERVER, DEFAULT_MAX_CONNECTIONS_WORKER, }; #[cfg(feature = "operator")] pub const DEFAULT_MAX_CONNECTIONS_OPERATOR: u32 = 2; pub async fn initial_connection() -> Result, error::Error> { let connect_options = get_database_url().await?.connect_options().await?; sqlx::postgres::PgPoolOptions::new() .max_connections(2) .connect_with(connect_options) .await .map_err(|err| Error::ConnectingToDatabase(err.to_string())) } /// Connect to the database for the Kubernetes operator process. /// /// Long-running operator pods need IAM RDS / Entra ID token refresh just like the server, /// otherwise new pool connections start failing once the initial token expires (~15 min). #[cfg(feature = "operator")] pub async fn operator_connection( #[cfg(all(feature = "enterprise", feature = "private"))] killpill_rx: tokio::sync::broadcast::Receiver<()>, ) -> anyhow::Result> { let database_url = get_database_url().await?; let pool = connect( database_url.clone(), DEFAULT_MAX_CONNECTIONS_OPERATOR, false, ) .await?; #[cfg(all(feature = "enterprise", feature = "private"))] spawn_token_refresh_task(pool.clone(), database_url, killpill_rx); Ok(pool) } pub async fn connect_db( server_mode: bool, indexer_mode: bool, worker_mode: bool, num_workers: i32, #[cfg(feature = "private")] killpill_rx: tokio::sync::broadcast::Receiver<()>, ) -> anyhow::Result> { use anyhow::Context; let database_url = get_database_url().await?; let max_connections = match std::env::var("DATABASE_CONNECTIONS") { Ok(n) => n.parse::().context("invalid DATABASE_CONNECTIONS")?, Err(_) => { if server_mode { DEFAULT_MAX_CONNECTIONS_SERVER } else if indexer_mode { DEFAULT_MAX_CONNECTIONS_INDEXER } else { DEFAULT_MAX_CONNECTIONS_WORKER + (num_workers.max(1) as u32) - 1 } } }; let pool = connect(database_url.clone(), max_connections, worker_mode).await?; #[cfg(all(feature = "enterprise", feature = "private"))] spawn_token_refresh_task(pool.clone(), database_url, killpill_rx); Ok(pool) } /// Spawn a background task that refreshes IAM RDS / Entra ID tokens before they expire /// and updates the pool's connect options so new connections use the fresh token. /// No-op for static (password-based) database URLs. #[cfg(all(feature = "enterprise", feature = "private"))] pub fn spawn_token_refresh_task( pool: sqlx::Pool, database_url: DatabaseUrl, mut killpill_rx: tokio::sync::broadcast::Receiver<()>, ) { let label = match &database_url { DatabaseUrl::IamRds(_) => "IAM RDS", DatabaseUrl::EntraId(_) => "Entra ID", DatabaseUrl::Static(_) => return, }; tokio::spawn(async move { loop { tokio::select! { _ = killpill_rx.recv() => { break; } _ = tokio::time::sleep(std::time::Duration::from_secs(10)) => { if !database_url.needs_refresh().await { continue; } let new_url = tokio::time::timeout( std::time::Duration::from_secs(10), get_database_url(), ) .await; match new_url { Ok(Ok(new_url)) => { match new_url.connect_options().await { Ok(connect_options) => { pool.set_connect_options(connect_options); tracing::info!("Refreshed {label} URL successfully"); } Err(e) => { tracing::error!( "Error getting {label} connect options, retrying in 10s: {e}" ); continue; } } } Ok(Err(e)) => { tracing::error!( "Error refreshing {label} URL, trying again in 10s: {e}" ); continue; } Err(e) => { tracing::error!( "Timeout after 10s refreshing {label} URL, trying again in 10s: {e}" ); continue; } } } } } }); } pub async fn connect( database_url: DatabaseUrl, max_connections: u32, worker_mode: bool, ) -> Result, error::Error> { use sqlx::Executor; use std::time::Duration; let mut pool_options = sqlx::postgres::PgPoolOptions::new() .min_connections(0) .max_connections(max_connections) .max_lifetime(Duration::from_secs(30 * 60)); // 30 mins if worker_mode { pool_options = pool_options.idle_timeout(Duration::from_secs(60)); } pool_options // Clears transaction state sqlx is not tracking, so a session left inside a // transaction is cleaned before the connection is handed to a borrower. See // `windmill_common::db::connection_reset` for how such a session comes about and why // this only runs once one has been observed. .before_acquire(|conn, _| { Box::pin(windmill_common::db::connection_reset::reset_before_acquire( conn, )) }) .after_connect(move |conn, _| { if worker_mode { Box::pin(async move { if let Err(e) = conn .execute( r#" SET enable_seqscan = OFF; SET statement_timeout = '5min'; SET idle_in_transaction_session_timeout = '10min'; SET tcp_keepalives_idle = 300; SET tcp_keepalives_interval = 60; SET tcp_keepalives_count = 10;"#, ) .await { tracing::error!("Error setting postgres settings: {}", e); } Ok(()) }) } else { Box::pin(async move { if let Err(e) = conn .execute( r#" SET statement_timeout = '5min'; SET idle_in_transaction_session_timeout = '10min'; SET tcp_keepalives_idle = 300; SET tcp_keepalives_interval = 60; SET tcp_keepalives_count = 10;"#, ) .await { tracing::error!("Error setting postgres settings: {}", e); } Ok(()) }) } }) .connect_with( database_url .connect_options() .await? .statement_cache_capacity(400), ) .await .map_err(|err| Error::ConnectingToDatabase(err.to_string())) }