use std::str::FromStr;
use crate::{
config::{DatabaseConfig, DatabaseType},
error::Error,
};
pub struct DatabaseManager {
pool: sqlx::Pool<sqlx::Any>,
db_type: DatabaseType,
}
impl DatabaseManager {
pub async fn new(db_config: DatabaseConfig) -> Result<Self, Error> {
let db_type = DatabaseType::from_str(&db_config.db_type)?;
let connect_str = match db_type {
DatabaseType::Sqlite => {
format!("sqlite://{}", db_config.database)
}
DatabaseType::Postgres => {
let host = db_config.host.unwrap_or("localhost".to_string());
let port = db_config.port.unwrap_or(5432);
let username = db_config.username.unwrap_or("postgres".to_string());
let password = db_config.password.unwrap_or("password".to_string());
format!(
"postgres://{}:{}@{}:{}/{}",
username, password, host, port, db_config.database
)
}
DatabaseType::MySql => {
let host = db_config.host.unwrap_or("localhost".to_string());
let port = db_config.port.unwrap_or(3306);
let username = db_config.username.unwrap_or("root".to_string());
let password = db_config.password.unwrap_or("password".to_string());
format!(
"mysql://{}:{}@{}:{}/{}",
username, password, host, port, db_config.database
)
}
};
let max_connections = db_config.max_connections.unwrap_or(12);
sqlx::any::install_default_drivers();
let pool = sqlx::any::AnyPoolOptions::new()
.max_connections(max_connections)
.connect(&connect_str)
.await?;
Ok(Self { pool, db_type })
}
pub fn pool(&self) -> &sqlx::Pool<sqlx::Any> {
&self.pool
}
pub fn db_type(&self) -> &DatabaseType {
&self.db_type
}
}
#[cfg(test)]
mod tests {
use crate::config::AppConfig;
use super::*;
#[tokio::test]
#[ignore = "Requires a database to be set up"]
async fn test_database_manager() {
let config = AppConfig::init().expect("Failed to load config");
let db_manager = DatabaseManager::new(config.botcat_capoo.database)
.await
.expect("Failed to create DatabaseManager");
let pool = db_manager.pool();
let db_type = db_manager.db_type();
println!("Database Type: {:?}", db_type);
let row: (i32,) = sqlx::query_as("SELECT 1")
.fetch_one(pool)
.await
.expect("Failed to execute query");
assert_eq!(row.0, 1);
}
}