use saddle_admission::DbCreditProfile;
use saddle_core::{ComponentLifecycle, LifecycleFuture};
use saddle_observability::Observer;
use saddle_runtime::startup_assembly::{StartupDbPoolFactory, StartupDbPoolOwner};
use crate::{Database, DatabaseConfig, SaddleError};
#[doc(hidden)]
pub struct StartupManagedDatabaseFactory {
config: Option<DatabaseConfig>,
observer: Option<Observer>,
}
impl StartupManagedDatabaseFactory {
pub fn new(config: Option<DatabaseConfig>, observer: Observer) -> Self {
Self {
config,
observer: Some(observer),
}
}
#[doc(hidden)]
pub fn awaiting_started_observability(config: Option<DatabaseConfig>) -> Self {
Self {
config,
observer: None,
}
}
}
#[doc(hidden)]
pub struct StartupManagedDatabaseOwner {
database: Option<Database>,
profile: DbCreditProfile,
}
#[doc(hidden)]
pub struct StartupManagedDatabaseBootstrap<'a> {
database: Option<&'a Database>,
}
impl StartupManagedDatabaseBootstrap<'_> {
pub fn existing(self) -> Option<Database> {
self.database.cloned()
}
}
impl StartupManagedDatabaseOwner {
pub fn bootstrap<R>(
self,
consume: impl for<'a> FnOnce(StartupManagedDatabaseBootstrap<'a>) -> R,
) -> (Self, R) {
let result = consume(StartupManagedDatabaseBootstrap {
database: self.database.as_ref(),
});
(self, result)
}
}
const _: () = ();
impl Drop for StartupManagedDatabaseOwner {
fn drop(&mut self) {
drop(self.database.take());
}
}
impl StartupDbPoolOwner for StartupManagedDatabaseOwner {
fn connection_capacity(&self) -> usize {
self.profile.connections
}
fn operation_capacity(&self) -> usize {
self.profile.operations
}
}
impl ComponentLifecycle for StartupManagedDatabaseOwner {
fn name(&self) -> &'static str {
"database"
}
fn start(&self) -> LifecycleFuture<'_> {
Box::pin(std::future::ready(Ok(())))
}
fn shutdown(&self) -> LifecycleFuture<'_> {
Box::pin(async move {
match self.database.as_ref() {
Some(database) => database.close().await,
None => Ok(()),
}
})
}
}
#[derive(Debug)]
#[doc(hidden)]
pub enum StartupManagedDatabaseError {
ConfigurationMismatch,
CapacityOverflow,
Connect(SaddleError),
}
impl StartupDbPoolFactory for StartupManagedDatabaseFactory {
type Owner = StartupManagedDatabaseOwner;
type Error = StartupManagedDatabaseError;
async fn construct(self, required: DbCreditProfile) -> Result<Self::Owner, Self::Error> {
if required.connections != required.operations {
return Err(StartupManagedDatabaseError::ConfigurationMismatch);
}
if required.connections == 0 {
return if self.config.is_none() {
Ok(StartupManagedDatabaseOwner {
database: None,
profile: required,
})
} else {
Err(StartupManagedDatabaseError::ConfigurationMismatch)
};
}
let config = self
.config
.ok_or(StartupManagedDatabaseError::ConfigurationMismatch)?;
let connections = u32::try_from(required.connections)
.map_err(|_| StartupManagedDatabaseError::CapacityOverflow)?;
let config = config.verified_connections(connections);
let observer = self
.observer
.ok_or(StartupManagedDatabaseError::ConfigurationMismatch)?;
let database = Database::connect(config, observer)
.await
.map_err(StartupManagedDatabaseError::Connect)?;
Ok(StartupManagedDatabaseOwner {
database: Some(database),
profile: required,
})
}
}
#[cfg(test)]
mod tests {
use std::{env, io, time::Duration};
use saddle_observability::ObserverConfig;
use sqlx::{Connection, mysql::MySqlConnection};
use super::*;
fn runtime() -> tokio::runtime::Runtime {
tokio::runtime::Builder::new_multi_thread()
.worker_threads(1)
.enable_all()
.build()
.unwrap()
}
fn observer() -> Observer {
Observer::with_writer(ObserverConfig::default(), io::sink()).unwrap()
}
#[test]
fn zero_profile_requires_no_database_configuration() {
let owner = runtime()
.block_on(
StartupManagedDatabaseFactory::new(None, observer()).construct(DbCreditProfile {
connections: 0,
operations: 0,
}),
)
.unwrap();
assert!(owner.database.is_none());
assert_eq!(owner.connection_capacity(), 0);
assert_eq!(owner.operation_capacity(), 0);
let (owner, database) = owner.bootstrap(|database| database.existing());
assert!(database.is_none());
drop(owner);
let result = runtime().block_on(
StartupManagedDatabaseFactory::new(
Some(DatabaseConfig::new("mysql://localhost/unused")),
observer(),
)
.construct(DbCreditProfile {
connections: 0,
operations: 0,
}),
);
assert!(matches!(
result,
Err(StartupManagedDatabaseError::ConfigurationMismatch)
));
}
#[test]
fn asymmetric_and_overflowing_profiles_fail_before_connect() {
let asymmetric = runtime().block_on(
StartupManagedDatabaseFactory::new(
Some(DatabaseConfig::new("mysql://localhost/unused")),
observer(),
)
.construct(DbCreditProfile {
connections: 1,
operations: 2,
}),
);
assert!(matches!(
asymmetric,
Err(StartupManagedDatabaseError::ConfigurationMismatch)
));
if usize::BITS > u32::BITS {
let overflow = runtime().block_on(
StartupManagedDatabaseFactory::new(
Some(DatabaseConfig::new("mysql://localhost/unused")),
observer(),
)
.construct(DbCreditProfile {
connections: u32::MAX as usize + 1,
operations: u32::MAX as usize + 1,
}),
);
assert!(matches!(
overflow,
Err(StartupManagedDatabaseError::CapacityOverflow)
));
}
}
#[test]
fn connect_failure_returns_without_a_physical_owner() {
let result = runtime().block_on(
StartupManagedDatabaseFactory::new(
Some(
DatabaseConfig::new("mysql://root@127.0.0.1:1/unreachable")
.acquire_timeout(Duration::from_millis(50)),
),
observer(),
)
.construct(DbCreditProfile {
connections: 1,
operations: 1,
}),
);
assert!(matches!(
result,
Err(StartupManagedDatabaseError::Connect(_))
));
}
#[test]
fn real_mariadb_constructs_exact_single_pool_and_drop_releases_owner() {
let Ok(url) = env::var("SADDLE_TEST_DATABASE_URL") else {
eprintln!("skipping startup pool adapter: SADDLE_TEST_DATABASE_URL is not set");
return;
};
let runtime = runtime();
let mut admin = runtime.block_on(MySqlConnection::connect(&url)).unwrap();
let baseline = runtime
.block_on(
sqlx::query_scalar::<_, i64>(
"SELECT COUNT(*) FROM information_schema.PROCESSLIST WHERE DB = DATABASE()",
)
.fetch_one(&mut admin),
)
.unwrap();
let owner = runtime
.block_on(
StartupManagedDatabaseFactory::new(
Some(DatabaseConfig::new(&url).max_connections(99)),
observer(),
)
.construct(DbCreditProfile {
connections: 2,
operations: 2,
}),
)
.unwrap();
assert_eq!(owner.connection_capacity(), 2);
assert_eq!(owner.operation_capacity(), 2);
let database = owner.database.as_ref().unwrap();
assert_eq!(database.pool.size(), 2);
assert_eq!(database.pool.num_idle(), 2);
let (owner, bootstrap_database) = owner.bootstrap(|database| database.existing());
let bootstrap_database = bootstrap_database.unwrap();
assert_eq!(bootstrap_database.pool.size(), 2);
drop(bootstrap_database);
let with_pool = runtime
.block_on(
sqlx::query_scalar::<_, i64>(
"SELECT COUNT(*) FROM information_schema.PROCESSLIST WHERE DB = DATABASE()",
)
.fetch_one(&mut admin),
)
.unwrap();
assert_eq!(with_pool, baseline + 2);
drop(owner);
runtime.block_on(async {
for _ in 0..50 {
let current = sqlx::query_scalar::<_, i64>(
"SELECT COUNT(*) FROM information_schema.PROCESSLIST WHERE DB = DATABASE()",
)
.fetch_one(&mut admin)
.await
.unwrap();
if current == baseline {
return;
}
tokio::time::sleep(Duration::from_millis(20)).await;
}
panic!("dropping the unique startup owner did not close its physical pool");
});
let lost = runtime
.block_on(
StartupManagedDatabaseFactory::new(Some(DatabaseConfig::new(&url)), observer())
.construct(DbCreditProfile {
connections: 2,
operations: 2,
}),
)
.unwrap();
let panic = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let _ = lost.bootstrap::<()>(|_| panic!("lost bootstrap consumer"));
}));
assert!(panic.is_err());
runtime.block_on(async {
for _ in 0..50 {
let current = sqlx::query_scalar::<_, i64>(
"SELECT COUNT(*) FROM information_schema.PROCESSLIST WHERE DB = DATABASE()",
)
.fetch_one(&mut admin)
.await
.unwrap();
if current == baseline {
return;
}
tokio::time::sleep(Duration::from_millis(20)).await;
}
panic!("lost bootstrap capability did not drop its physical owner");
});
runtime.block_on(admin.close()).unwrap();
runtime.shutdown_background();
}
}