use std::{mem::take, time::Duration};
use super::Database;
use crate::{
ON_CONNECT_FILE, ON_RESET_FILE,
app_config::AppConfig,
webserver::database::{DbInfo, SupportedDatabase},
};
use anyhow::Context;
use futures_util::future::BoxFuture;
use sqlx::connection::{ConnectOptions, Connection};
use sqlx::executor::Executor;
use sqlx::odbc::OdbcConnectOptions;
use sqlx::{
any::{Any, AnyConnectOptions, AnyConnection, AnyKind},
pool::PoolOptions,
sqlite::{Function, SqliteConnectOptions, SqliteFunctionCtx},
};
impl Database {
pub async fn init(config: &AppConfig) -> anyhow::Result<Self> {
let database_url = &config.database_url;
let mut connect_options: AnyConnectOptions = database_url
.parse()
.with_context(|| format!("\"{database_url}\" is not a valid database URL. Please change the \"database_url\" option in the configuration file."))?;
if let Some(password) = &config.database_password {
set_database_password(&mut connect_options, password);
}
connect_options.log_statements(log::LevelFilter::Trace);
connect_options.log_slow_statements(log::LevelFilter::Warn, Duration::from_millis(250));
log::debug!(
"Connecting to a {:?} database on {}",
connect_options.kind(),
database_url
);
set_custom_connect_options(&mut connect_options, config);
log::debug!("Connecting to database: {database_url}");
let mut retries = config.database_connection_retries;
let mut conn: AnyConnection = loop {
match AnyConnection::connect_with(&connect_options).await {
Ok(c) => break c,
Err(e) => {
if retries == 0 {
return Err(anyhow::Error::new(e)
.context(format!("Unable to open connection to {database_url}")));
}
log::warn!("Failed to connect to the database: {e:#}. Retrying in 5 seconds.");
retries -= 1;
tokio::time::sleep(Duration::from_secs(5)).await;
}
}
};
let dbms_name: String = conn.dbms_name().await?;
let database_type = SupportedDatabase::from_dbms_name(&dbms_name);
drop(conn);
let db_kind = connect_options.kind();
let pool = Self::create_pool_options(config, db_kind)
.connect_with(connect_options)
.await
.with_context(|| format!("Unable to open connection pool to {database_url}"))?;
log::debug!("Initialized {dbms_name:?} database pool: {pool:#?}");
Ok(Database {
connection: pool,
info: DbInfo {
dbms_name,
database_type,
kind: db_kind,
},
})
}
fn create_pool_options(config: &AppConfig, kind: AnyKind) -> PoolOptions<Any> {
let mut pool_options = PoolOptions::new()
.max_connections(if let Some(max) = config.max_database_pool_connections {
max
} else {
match kind {
AnyKind::Postgres | AnyKind::Odbc => 50, AnyKind::MySql => 75,
AnyKind::Sqlite => {
if config.database_url.contains(":memory:") {
128
} else {
16
}
}
AnyKind::Mssql => 100,
}
})
.idle_timeout(config.database_connection_idle_timeout)
.max_lifetime(config.database_connection_max_lifetime)
.acquire_timeout(Duration::from_secs_f64(
config.database_connection_acquire_timeout_seconds,
));
pool_options = add_on_return_to_pool(config, pool_options);
pool_options = add_on_connection_handler(config, pool_options);
pool_options
}
}
fn add_on_return_to_pool(config: &AppConfig, pool_options: PoolOptions<Any>) -> PoolOptions<Any> {
let on_disconnect_file = config.configuration_directory.join(ON_RESET_FILE);
let sql = if on_disconnect_file.exists() {
log::info!(
"Creating a custom SQL connection cleanup handler from {}",
on_disconnect_file.display()
);
match std::fs::read_to_string(&on_disconnect_file) {
Ok(sql) => {
log::trace!("The custom SQL connection cleanup handler is:\n{sql}");
Some(std::sync::Arc::new(sql))
}
Err(e) => {
log::error!(
"Unable to read the file {}: {}",
on_disconnect_file.display(),
e
);
None
}
}
} else {
log::debug!(
"Not creating a custom SQL connection cleanup handler because {} does not exist",
on_disconnect_file.display()
);
None
};
pool_options.after_release(move |conn, meta| {
let sql = sql.clone();
Box::pin(async move {
if let Some(sql) = sql {
on_return_to_pool(conn, meta, sql).await
} else {
Ok(true)
}
})
})
}
fn on_return_to_pool(
conn: &mut AnyConnection,
meta: sqlx::pool::PoolConnectionMetadata,
sql: std::sync::Arc<String>,
) -> BoxFuture<'_, Result<bool, sqlx::error::Error>> {
use sqlx::row::Row;
Box::pin(async move {
log::trace!("Running the custom SQL connection cleanup handler. {meta:?}");
let query_result = conn.fetch_optional(sql.as_str()).await?;
if let Some(query_result) = query_result {
let is_healthy = query_result.try_get::<bool, _>(0);
log::debug!("Is the connection healthy? {is_healthy:?}");
is_healthy
} else {
log::debug!("No result from the custom SQL connection cleanup handler");
Ok(true)
}
})
}
fn add_on_connection_handler(
config: &AppConfig,
pool_options: PoolOptions<Any>,
) -> PoolOptions<Any> {
let on_connect_file = config.configuration_directory.join(ON_CONNECT_FILE);
let on_connect_file_display = on_connect_file.display().to_string();
let sql = if on_connect_file.exists() {
log::info!(
"Creating a custom SQL database connection handler from {}",
on_connect_file.display()
);
match std::fs::read_to_string(&on_connect_file) {
Ok(sql) => {
log::trace!("The custom SQL database connection handler is:\n{sql}");
Some(std::sync::Arc::new(sql))
}
Err(e) => {
log::error!(
"Unable to read the file {}: {}",
on_connect_file.display(),
e
);
None
}
}
} else {
log::debug!(
"Not creating a custom SQL database connection handler because {} does not exist",
on_connect_file.display()
);
None
};
pool_options.after_connect(move |conn, _| {
let sql = sql.clone();
let on_connect_file_display = on_connect_file_display.clone();
Box::pin(async move {
if let Some(sql) = sql {
log::debug!("Running {on_connect_file_display} on new connection");
let r = conn.execute(sql.as_str()).await?;
log::debug!("Finished running connection handler on new connection: {r:?}");
}
Ok(())
})
})
}
fn set_custom_connect_options(options: &mut AnyConnectOptions, config: &AppConfig) {
if let Some(sqlite_options) = options.as_sqlite_mut() {
set_custom_connect_options_sqlite(sqlite_options, config);
}
if let Some(odbc_options) = options.as_odbc_mut() {
set_custom_connect_options_odbc(odbc_options, config);
}
}
fn set_custom_connect_options_sqlite(
sqlite_options: &mut SqliteConnectOptions,
config: &AppConfig,
) {
for extension_name in &config.sqlite_extensions {
log::info!("Loading SQLite extension: {extension_name}");
*sqlite_options = take(sqlite_options).extension(extension_name.clone());
}
*sqlite_options = take(sqlite_options)
.collation("NOCASE", |a, b| a.to_lowercase().cmp(&b.to_lowercase()))
.function(make_sqlite_fun("upper", str::to_uppercase))
.function(make_sqlite_fun("lower", str::to_lowercase));
}
fn make_sqlite_fun(name: &str, f: fn(&str) -> String) -> Function {
Function::new(name, move |ctx: &SqliteFunctionCtx| {
let arg = ctx.try_get_arg::<Option<&str>>(0);
match arg {
Ok(Some(s)) => ctx.set_result(f(s)),
Ok(None) => ctx.set_result(None::<String>),
Err(e) => ctx.set_error(&e.to_string()),
}
})
}
fn set_custom_connect_options_odbc(odbc_options: &mut OdbcConnectOptions, config: &AppConfig) {
let batch_size = config.max_pending_rows.clamp(1, 1024);
odbc_options.batch_size(batch_size);
log::trace!("ODBC batch size set to {batch_size}");
odbc_options.max_column_size(None);
}
fn set_database_password(options: &mut AnyConnectOptions, password: &str) {
if let Some(opts) = options.as_postgres_mut() {
*opts = take(opts).password(password);
} else if let Some(opts) = options.as_mysql_mut() {
*opts = take(opts).password(password);
} else if let Some(opts) = options.as_mssql_mut() {
*opts = take(opts).password(password);
} else if let Some(_opts) = options.as_odbc_mut() {
log::warn!(
"Setting a password for an ODBC connection is not supported via separate config; include credentials in the DSN or connection string"
);
} else if let Some(_opts) = options.as_sqlite_mut() {
log::warn!("Setting a password for a SQLite database is not supported");
} else {
unreachable!("Unsupported database type");
}
}
#[cfg(test)]
mod tests {
use super::*;
use sqlx::row::Row;
use tempfile::TempDir;
fn config_in(dir: &TempDir, database_url: &str) -> AppConfig {
let mut config = crate::app_config::tests::test_config();
config.database_url = database_url.to_owned();
config.configuration_directory = dir.path().to_path_buf();
config.max_database_pool_connections = Some(1);
config
}
#[actix_web::test]
async fn on_connect_sql_runs_on_new_pool_connections() {
let dir = TempDir::new().unwrap();
std::fs::write(
dir.path().join(ON_CONNECT_FILE),
"CREATE TEMPORARY TABLE on_connect_marker(value TEXT);
INSERT INTO on_connect_marker(value) VALUES ('on_connect ran');",
)
.unwrap();
let db = Database::init(&config_in(&dir, "sqlite::memory:"))
.await
.unwrap();
let value: String = sqlx::query::query("SELECT value FROM on_connect_marker")
.fetch_one(&db.connection)
.await
.unwrap()
.try_get(0)
.unwrap();
assert_eq!(value, "on_connect ran");
}
#[actix_web::test]
async fn on_reset_sql_runs_when_a_connection_returns_to_the_pool() {
let dir = TempDir::new().unwrap();
std::fs::write(
dir.path().join(ON_CONNECT_FILE),
"CREATE TABLE IF NOT EXISTS reset_log(id INTEGER);",
)
.unwrap();
std::fs::write(
dir.path().join(ON_RESET_FILE),
"INSERT INTO reset_log(id) VALUES (1);
SELECT 1 AS is_healthy;",
)
.unwrap();
let db_file = dir.path().join("on_reset.db");
let config = config_in(&dir, &format!("sqlite://{}?mode=rwc", db_file.display()));
let db = Database::init(&config).await.unwrap();
drop(db.connection.acquire().await.unwrap());
let resets: i64 = sqlx::query::query("SELECT COUNT(*) FROM reset_log")
.fetch_one(&db.connection)
.await
.unwrap()
.try_get(0)
.unwrap();
assert!(resets > 0);
}
#[tokio::test(start_paused = true)]
async fn connection_retries_wait_five_seconds_before_giving_up() {
for (retries, expected_wait) in [(0, Duration::ZERO), (2, Duration::from_secs(10))] {
let dir = TempDir::new().unwrap();
let missing = dir.path().join("nonexistent_directory").join("db.sqlite");
let mut config = config_in(&dir, &format!("sqlite://{}", missing.display()));
config.database_connection_retries = retries;
let start = tokio::time::Instant::now();
let Err(error) = Database::init(&config).await else {
panic!("connecting to a missing database must fail");
};
assert_eq!(start.elapsed(), expected_wait, "{retries} retries");
assert!(
format!("{error:#}").contains("Unable to open connection to"),
"{error:#}"
);
}
}
}