From a91d8badbffa7aab748cb311047fc14a30151a97 Mon Sep 17 00:00:00 2001 From: Richard Russell <2265225+rars@users.noreply.github.com> Date: Thu, 2 Jul 2026 14:15:51 +0100 Subject: [PATCH] refactor: use r2d2 connection pool for DB connections --- src-tauri/Cargo.lock | 21 ++++++++ src-tauri/Cargo.toml | 2 +- src-tauri/src/commands/app.rs | 5 +- src-tauri/src/commands/electricity.rs | 24 ++++----- src-tauri/src/commands/gas.rs | 24 ++++----- src-tauri/src/commands/mod.rs | 2 + src-tauri/src/commands/profiles.rs | 4 +- src-tauri/src/data/consumption.rs | 34 ++++++------ src-tauri/src/data/energy_profile.rs | 27 +++++----- src-tauri/src/data/mod.rs | 4 +- src-tauri/src/data/tariff.rs | 75 +++++++++++++-------------- src-tauri/src/db.rs | 15 ++---- src-tauri/src/download.rs | 44 +++++++--------- src-tauri/src/main.rs | 30 ++++++++--- src-tauri/src/utils.rs | 8 ++- 15 files changed, 169 insertions(+), 150 deletions(-) diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 76123c9..aae119e 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -1085,6 +1085,7 @@ dependencies = [ "diesel_derives", "downcast-rs", "libsqlite3-sys", + "r2d2", "sqlite-wasm-rs", "time", ] @@ -3706,6 +3707,17 @@ version = "6.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" +[[package]] +name = "r2d2" +version = "0.8.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "51de85fb3fb6524929c8a2eb85e6b6d363de4e8c48f9e2c2eac4944abc181c93" +dependencies = [ + "log", + "parking_lot", + "scheduled-thread-pool", +] + [[package]] name = "radium" version = "0.7.0" @@ -4114,6 +4126,15 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "scheduled-thread-pool" +version = "0.2.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3cbc66816425a074528352f5789333ecff06ca41b36b0b0efdfbb29edc391a19" +dependencies = [ + "parking_lot", +] + [[package]] name = "schemars" version = "0.8.22" diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 516f22c..d4c213b 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -16,7 +16,7 @@ tauri-build = { version = "2.0.0", features = [] } [dependencies] chrono = { version = "0.4", features = ["serde"] } -diesel = { version = "2.0.0", features = ["sqlite", "chrono"] } +diesel = { version = "2.0.0", features = ["sqlite", "chrono", "r2d2"] } diesel_migrations = "2.0" git-version = "0.3.9" glowmarkt = { version = "0.5.3" } diff --git a/src-tauri/src/commands/app.rs b/src-tauri/src/commands/app.rs index 3f16ec6..e5fec10 100644 --- a/src-tauri/src/commands/app.rs +++ b/src-tauri/src/commands/app.rs @@ -129,10 +129,7 @@ pub async fn reset(app_handle: AppHandle, app_state: State<'_, AppState>) -> Res } fn reset_database(app_state: &AppState) -> Result<(), ApiError> { - let mut conn = app_state - .db - .lock() - .map_err(|_| ApiError::MutexPoisonedError { name: "db".into() })?; + let mut conn = app_state.db_pool.get()?; revert_all_migrations(&mut conn); db::run_migrations(&mut conn); diff --git a/src-tauri/src/commands/electricity.rs b/src-tauri/src/commands/electricity.rs index 5780036..16939e8 100644 --- a/src-tauri/src/commands/electricity.rs +++ b/src-tauri/src/commands/electricity.rs @@ -50,10 +50,10 @@ pub async fn get_raw_electricity_consumption( let start = parse_iso_string_to_naive_date(&start_date)?; let end = parse_iso_string_to_naive_date(&end_date)?; - let db_connection_clone = app_state.db.clone(); + let connection_pool_clone = app_state.db_pool.clone(); let result = async_runtime::spawn_blocking(move || { - let repository = SqliteElectricityConsumptionRepository::new(db_connection_clone); + let repository = SqliteElectricityConsumptionRepository::new(connection_pool_clone); repository.get_raw(start, end) }) @@ -93,10 +93,10 @@ pub async fn get_daily_electricity_consumption( let start = parse_iso_string_to_naive_date(&start_date)?; let end = parse_iso_string_to_naive_date(&end_date)?; - let db_connection_clone = app_state.db.clone(); + let connection_pool_clone = app_state.db_pool.clone(); let daily_consumption = async_runtime::spawn_blocking(move || { - let repository = SqliteElectricityConsumptionRepository::new(db_connection_clone); + let repository = SqliteElectricityConsumptionRepository::new(connection_pool_clone); repository.get_daily(start, end) }) @@ -127,10 +127,10 @@ pub async fn get_monthly_electricity_consumption( let start = parse_iso_string_to_naive_date(&start_date)?; let end = parse_iso_string_to_naive_date(&end_date)?; - let db_connection_clone = app_state.db.clone(); + let connection_pool_clone = app_state.db_pool.clone(); let monthly_consumption = async_runtime::spawn_blocking(move || { - let repository = SqliteElectricityConsumptionRepository::new(db_connection_clone); + let repository = SqliteElectricityConsumptionRepository::new(connection_pool_clone); repository.get_monthly(start, end) }) @@ -151,10 +151,10 @@ pub async fn get_monthly_electricity_consumption( pub async fn get_electricity_tariff_history( app_state: State<'_, AppState>, ) -> Result { - let db_connection_clone = app_state.db.clone(); + let connection_pool_clone = app_state.db_pool.clone(); let standing_charge_history = async_runtime::spawn_blocking(move || { - let repository = SqliteElectricityTariffRepository::new(db_connection_clone); + let repository = SqliteElectricityTariffRepository::new(connection_pool_clone); repository.get_standing_charge_history() }) @@ -162,10 +162,10 @@ pub async fn get_electricity_tariff_history( .map_err(|e| ApiError::Custom(format!("Error: {}", e)))? .map_err(ApiError::RepositoryError)?; - let db_connection_clone = app_state.db.clone(); + let connection_pool_clone = app_state.db_pool.clone(); let unit_price_history = async_runtime::spawn_blocking(move || { - let repository = SqliteElectricityTariffRepository::new(db_connection_clone); + let repository = SqliteElectricityTariffRepository::new(connection_pool_clone); repository.get_unit_price_history() }) @@ -206,10 +206,10 @@ pub async fn get_electricity_cost_history( let start = parse_iso_string_to_naive_date(&start_date)?; let end = parse_iso_string_to_naive_date(&end_date)?; - let db_connection_clone = app_state.db.clone(); + let connection_pool_clone = app_state.db_pool.clone(); let mut consumption = async_runtime::spawn_blocking(move || { - let repository = SqliteElectricityConsumptionRepository::new(db_connection_clone); + let repository = SqliteElectricityConsumptionRepository::new(connection_pool_clone); repository.get_daily(start, end) }) diff --git a/src-tauri/src/commands/gas.rs b/src-tauri/src/commands/gas.rs index 3401c44..9d1192e 100644 --- a/src-tauri/src/commands/gas.rs +++ b/src-tauri/src/commands/gas.rs @@ -46,10 +46,10 @@ pub async fn get_raw_gas_consumption( let start = parse_iso_string_to_naive_date(&start_date)?; let end = parse_iso_string_to_naive_date(&end_date)?; - let db_connection_clone = app_state.db.clone(); + let connection_pool_clone = app_state.db_pool.clone(); let raw_consumption = async_runtime::spawn_blocking(move || { - let repository = SqliteGasConsumptionRepository::new(db_connection_clone); + let repository = SqliteGasConsumptionRepository::new(connection_pool_clone); repository.get_raw(start, end) }) @@ -77,10 +77,10 @@ pub async fn get_daily_gas_consumption( let start = parse_iso_string_to_naive_date(&start_date)?; let end = parse_iso_string_to_naive_date(&end_date)?; - let db_connection_clone = app_state.db.clone(); + let connection_pool_clone = app_state.db_pool.clone(); let daily_consumption = async_runtime::spawn_blocking(move || { - let repository = SqliteGasConsumptionRepository::new(db_connection_clone); + let repository = SqliteGasConsumptionRepository::new(connection_pool_clone); repository.get_daily(start, end) }) @@ -108,10 +108,10 @@ pub async fn get_monthly_gas_consumption( let start = parse_iso_string_to_naive_date(&start_date)?; let end = parse_iso_string_to_naive_date(&end_date)?; - let db_connection_clone = app_state.db.clone(); + let connection_pool_clone = app_state.db_pool.clone(); let monthly_consumption = async_runtime::spawn_blocking(move || { - let repository = SqliteGasConsumptionRepository::new(db_connection_clone); + let repository = SqliteGasConsumptionRepository::new(connection_pool_clone); repository.get_monthly(start, end) }) @@ -132,10 +132,10 @@ pub async fn get_monthly_gas_consumption( pub async fn get_gas_tariff_history( app_state: State<'_, AppState>, ) -> Result { - let db_connection_clone = app_state.db.clone(); + let connection_pool_clone = app_state.db_pool.clone(); let standing_charge_history = async_runtime::spawn_blocking(move || { - let repository = SqliteGasTariffRepository::new(db_connection_clone); + let repository = SqliteGasTariffRepository::new(connection_pool_clone); repository.get_standing_charge_history() }) @@ -143,10 +143,10 @@ pub async fn get_gas_tariff_history( .map_err(|e| ApiError::Custom(format!("Error: {}", e)))? .map_err(ApiError::RepositoryError)?; - let db_connection_clone = app_state.db.clone(); + let connection_pool_clone = app_state.db_pool.clone(); let unit_price_history = async_runtime::spawn_blocking(move || { - let repository = SqliteGasTariffRepository::new(db_connection_clone); + let repository = SqliteGasTariffRepository::new(connection_pool_clone); repository.get_unit_price_history() }) @@ -187,10 +187,10 @@ pub async fn get_gas_cost_history( let start = parse_iso_string_to_naive_date(&start_date)?; let end = parse_iso_string_to_naive_date(&end_date)?; - let db_connection_clone = app_state.db.clone(); + let connection_pool_clone = app_state.db_pool.clone(); let mut consumption = async_runtime::spawn_blocking(move || { - let repository = SqliteGasConsumptionRepository::new(db_connection_clone); + let repository = SqliteGasConsumptionRepository::new(connection_pool_clone); repository.get_daily(start, end) }) diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index 3278430..b573e16 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -26,6 +26,8 @@ pub enum ApiError { JoinError(#[from] tokio::task::JoinError), #[error("Mutex '{name}' is poisoned")] MutexPoisonedError { name: String }, + #[error("Error with DB connection pool: {0}")] + ConnectionPoolError(#[from] diesel::r2d2::PoolError), } impl serde::Serialize for ApiError { diff --git a/src-tauri/src/commands/profiles.rs b/src-tauri/src/commands/profiles.rs index b8d3da0..cd33776 100644 --- a/src-tauri/src/commands/profiles.rs +++ b/src-tauri/src/commands/profiles.rs @@ -21,7 +21,7 @@ pub struct EnergyProfileUpdateParam { #[tauri::command] pub fn get_energy_profiles(app_state: State<'_, AppState>) -> Result, ApiError> { - let repository = SqliteEnergyProfileRepository::new(app_state.db.clone()); + let repository = SqliteEnergyProfileRepository::new(app_state.db_pool.clone()); repository .get_all_energy_profiles() @@ -45,7 +45,7 @@ pub async fn update_energy_profile_settings( }) .collect(); - let repository = SqliteEnergyProfileRepository::new(app_state.db.clone()); + let repository = SqliteEnergyProfileRepository::new(app_state.db_pool.clone()); for (energy_profile_id, is_active, start) in update_settings? { debug!("Updating {}, {}, {}", energy_profile_id, is_active, start); diff --git a/src-tauri/src/data/consumption.rs b/src-tauri/src/data/consumption.rs index 7d4beff..19d65ed 100644 --- a/src-tauri/src/data/consumption.rs +++ b/src-tauri/src/data/consumption.rs @@ -1,14 +1,16 @@ -use std::sync::{Arc, Mutex, MutexGuard}; - use chrono::{NaiveDate, NaiveDateTime}; use diesel::dsl::sql; use diesel::insert_into; +use diesel::r2d2::ConnectionManager; +use diesel::r2d2::Pool; +use diesel::r2d2::PooledConnection; use diesel::SqliteConnection; use diesel::{prelude::*, upsert::excluded}; use log::error; use rust_decimal::prelude::ToPrimitive; use rust_decimal::Decimal; +use crate::db::SqliteConnectionPool; use crate::schema::{electricity_consumption, gas_consumption}; use crate::utils::london_date_id_to_naive_date; use crate::utils::{ @@ -84,18 +86,18 @@ pub trait ConsumptionRepository { } pub struct SqliteElectricityConsumptionRepository { - conn: Arc>, + connection_pool: Pool>, } impl SqliteElectricityConsumptionRepository { - pub fn new(conn: Arc>) -> Self { - Self { conn } + pub fn new(connection_pool: Pool>) -> Self { + Self { connection_pool } } - fn get_connection(&self) -> RepositoryResult> { - self.conn - .lock() - .map_err(|_| RepositoryError::SqliteConnectionMutexPoisonedError()) + fn get_connection( + &self, + ) -> RepositoryResult>> { + Ok(self.connection_pool.get()?) } } @@ -220,18 +222,18 @@ impl ConsumptionRepository>, + connection_pool: SqliteConnectionPool, } impl SqliteGasConsumptionRepository { - pub fn new(connection: Arc>) -> Self { - Self { connection } + pub fn new(connection_pool: SqliteConnectionPool) -> Self { + Self { connection_pool } } - fn get_connection(&self) -> RepositoryResult> { - self.connection - .lock() - .map_err(|_| RepositoryError::SqliteConnectionMutexPoisonedError()) + fn get_connection( + &self, + ) -> RepositoryResult>> { + Ok(self.connection_pool.get()?) } } diff --git a/src-tauri/src/data/energy_profile.rs b/src-tauri/src/data/energy_profile.rs index 609c3f2..4700bff 100644 --- a/src-tauri/src/data/energy_profile.rs +++ b/src-tauri/src/data/energy_profile.rs @@ -1,10 +1,11 @@ -use std::sync::{Arc, Mutex, MutexGuard}; - use chrono::{Datelike, Local, NaiveDateTime}; use diesel::dsl::*; use diesel::prelude::*; +use diesel::r2d2::ConnectionManager; +use diesel::r2d2::PooledConnection; use serde::Serialize; +use crate::db::SqliteConnectionPool; use crate::schema::energy_profile; #[derive(Serialize, Queryable, Debug)] @@ -47,24 +48,22 @@ pub trait EnergyProfileRepository { } pub struct SqliteEnergyProfileRepository { - conn: Arc>, + connection_pool: SqliteConnectionPool, } impl SqliteEnergyProfileRepository { - pub fn new(conn: Arc>) -> Self { - Self { conn } + pub fn new(connection_pool: SqliteConnectionPool) -> Self { + Self { connection_pool } } - fn connection(&self) -> MutexGuard<'_, SqliteConnection> { - self.conn - .lock() - .expect("Could not acquire lock on SqliteConnection") + fn get_connection(&self) -> PooledConnection> { + self.connection_pool.get().unwrap() } } impl EnergyProfileRepository for SqliteEnergyProfileRepository { fn get_energy_profile(&self, name: &str) -> QueryResult { - let mut conn = self.connection(); + let mut conn = self.get_connection(); energy_profile::table .filter(energy_profile::name.eq(name)) @@ -72,7 +71,7 @@ impl EnergyProfileRepository for SqliteEnergyProfileRepository { } fn get_all_energy_profiles(&self) -> QueryResult> { - let mut conn = self.connection(); + let mut conn = self.get_connection(); energy_profile::table.load::(&mut *conn) } @@ -89,7 +88,7 @@ impl EnergyProfileRepository for SqliteEnergyProfileRepository { base_unit, }; - let mut conn = self.connection(); + let mut conn = self.get_connection(); diesel::insert_into(energy_profile::table) .values(&new_profile) @@ -109,7 +108,7 @@ impl EnergyProfileRepository for SqliteEnergyProfileRepository { ) -> QueryResult { use crate::schema::energy_profile::dsl::*; - let mut conn = self.connection(); + let mut conn = self.get_connection(); diesel::update(energy_profile.find(energy_profile_id_param)) .set(( @@ -136,7 +135,7 @@ impl EnergyProfileRepository for SqliteEnergyProfileRepository { ) -> QueryResult { use crate::schema::energy_profile::dsl::*; - let mut conn = self.connection(); + let mut conn = self.get_connection(); diesel::update(energy_profile.find(energy_profile_id_param)) .set(( diff --git a/src-tauri/src/data/mod.rs b/src-tauri/src/data/mod.rs index 2c5bdd1..7a9a08a 100644 --- a/src-tauri/src/data/mod.rs +++ b/src-tauri/src/data/mod.rs @@ -6,6 +6,6 @@ pub mod tariff; pub enum RepositoryError { #[error("Database error: {0}")] DieselError(#[from] diesel::result::Error), - #[error("Mutex guarding SQLite connection is poisoned")] - SqliteConnectionMutexPoisonedError(), + #[error("Error with connection pool: {0}")] + ConnectionPoolError(#[from] diesel::r2d2::PoolError), } diff --git a/src-tauri/src/data/tariff.rs b/src-tauri/src/data/tariff.rs index 7ace779..8ade27d 100644 --- a/src-tauri/src/data/tariff.rs +++ b/src-tauri/src/data/tariff.rs @@ -1,6 +1,5 @@ -use std::sync::{Arc, Mutex, MutexGuard}; - use chrono::NaiveDateTime; +use diesel::r2d2::{ConnectionManager, Pool, PooledConnection}; use diesel::sql_types::{Double, Timestamp}; use diesel::{insert_into, sql_query, SqliteConnection}; use diesel::{prelude::*, upsert::excluded}; @@ -91,18 +90,18 @@ pub trait TariffRepository { } pub struct SqliteElectricityTariffRepository { - conn: Arc>, + connection_pool: Pool>, } impl SqliteElectricityTariffRepository { - pub fn new(conn: Arc>) -> Self { - Self { conn } + pub fn new(connection_pool: Pool>) -> Self { + Self { connection_pool } } - fn get_connection(&self) -> RepositoryResult> { - self.conn - .lock() - .map_err(|_| RepositoryError::SqliteConnectionMutexPoisonedError()) + fn get_connection( + &self, + ) -> RepositoryResult>> { + Ok(self.connection_pool.get()?) } } @@ -137,7 +136,7 @@ impl TariffRepository for SqliteElectricityTariffRepos let query = r#" WITH extracted_standing_charge AS ( - SELECT + SELECT effective_date, (SELECT json_extract(value, '$') FROM json_tree(electricity_tariff_plan.plan) @@ -146,14 +145,14 @@ impl TariffRepository for SqliteElectricityTariffRepos FROM electricity_tariff_plan ), price_changes AS ( - SELECT - effective_date, + SELECT + effective_date, standing_charge_pence, LAG(standing_charge_pence) OVER (ORDER BY effective_date) AS previous_price FROM extracted_standing_charge ) - SELECT - pc.effective_date AS start_date, + SELECT + pc.effective_date AS start_date, pc.standing_charge_pence FROM price_changes AS pc WHERE pc.standing_charge_pence <> COALESCE(pc.previous_price, pc.standing_charge_pence) @@ -169,7 +168,7 @@ impl TariffRepository for SqliteElectricityTariffRepos let query = r#" WITH extracted_unit_price AS ( - SELECT + SELECT effective_date, (SELECT json_extract(value, '$') FROM json_tree(electricity_tariff_plan.plan) @@ -178,17 +177,17 @@ impl TariffRepository for SqliteElectricityTariffRepos FROM electricity_tariff_plan ), price_changes AS ( - SELECT - effective_date, + SELECT + effective_date, unit_price_pence, LAG(unit_price_pence) OVER (ORDER BY effective_date) AS previous_price FROM extracted_unit_price ) - SELECT - pc.effective_date AS price_effective_time, + SELECT + pc.effective_date AS price_effective_time, pc.unit_price_pence FROM price_changes AS pc - WHERE pc.unit_price_pence <> COALESCE(pc.previous_price, pc.unit_price_pence) + WHERE pc.unit_price_pence <> COALESCE(pc.previous_price, pc.unit_price_pence) OR pc.previous_price IS NULL ORDER BY pc.effective_date; "#; @@ -301,18 +300,18 @@ impl TariffRepository for SqliteElectricityTariffRepository { */ pub struct SqliteGasTariffRepository { - conn: Arc>, + connection_pool: Pool>, } impl SqliteGasTariffRepository { - pub fn new(conn: Arc>) -> Self { - Self { conn } + pub fn new(connection_pool: Pool>) -> Self { + Self { connection_pool } } - fn get_connection(&self) -> RepositoryResult> { - self.conn - .lock() - .map_err(|_| RepositoryError::SqliteConnectionMutexPoisonedError()) + fn get_connection( + &self, + ) -> RepositoryResult>> { + Ok(self.connection_pool.get()?) } } @@ -346,7 +345,7 @@ impl TariffRepository for SqliteGasTariffRepository { let query = r#" WITH extracted_standing_charge AS ( - SELECT + SELECT effective_date, (SELECT json_extract(value, '$') FROM json_tree(gas_tariff_plan.plan) @@ -355,14 +354,14 @@ impl TariffRepository for SqliteGasTariffRepository { FROM gas_tariff_plan ), price_changes AS ( - SELECT - effective_date, + SELECT + effective_date, standing_charge_pence, LAG(standing_charge_pence) OVER (ORDER BY effective_date) AS previous_price FROM extracted_standing_charge ) - SELECT - pc.effective_date AS start_date, + SELECT + pc.effective_date AS start_date, pc.standing_charge_pence FROM price_changes AS pc WHERE pc.standing_charge_pence <> COALESCE(pc.previous_price, pc.standing_charge_pence) @@ -378,7 +377,7 @@ impl TariffRepository for SqliteGasTariffRepository { let query = r#" WITH extracted_unit_price AS ( - SELECT + SELECT effective_date, (SELECT json_extract(value, '$') FROM json_tree(gas_tariff_plan.plan) @@ -387,17 +386,17 @@ impl TariffRepository for SqliteGasTariffRepository { FROM gas_tariff_plan ), price_changes AS ( - SELECT - effective_date, + SELECT + effective_date, unit_price_pence, LAG(unit_price_pence) OVER (ORDER BY effective_date) AS previous_price FROM extracted_unit_price ) - SELECT - pc.effective_date AS price_effective_time, + SELECT + pc.effective_date AS price_effective_time, pc.unit_price_pence FROM price_changes AS pc - WHERE pc.unit_price_pence <> COALESCE(pc.previous_price, pc.unit_price_pence) + WHERE pc.unit_price_pence <> COALESCE(pc.previous_price, pc.unit_price_pence) OR pc.previous_price IS NULL ORDER BY pc.effective_date; "#; diff --git a/src-tauri/src/db.rs b/src-tauri/src/db.rs index 66073e3..82d3e1c 100644 --- a/src-tauri/src/db.rs +++ b/src-tauri/src/db.rs @@ -1,18 +1,16 @@ -use chrono_tz::Europe::London; use diesel::prelude::*; +use diesel::r2d2::{ConnectionManager, Pool}; use diesel::{sqlite::SqliteConnection, Connection, ExpressionMethods}; use diesel_migrations::{embed_migrations, EmbeddedMigrations, MigrationHarness}; use log::info; use crate::data::RepositoryError; use crate::utils::utc_timestamp_to_london_date_id; -use crate::{ - schema::electricity_consumption::{london_date_id, table}, - AppError, -}; const MIGRATIONS: EmbeddedMigrations = embed_migrations!("migrations"); +pub type SqliteConnectionPool = Pool>; + pub fn run_migrations(conn: &mut SqliteConnection) { conn.run_pending_migrations(MIGRATIONS) .expect("Error running migrations"); @@ -23,13 +21,6 @@ pub fn revert_all_migrations(conn: &mut SqliteConnection) { .expect("Error reverting migrations"); } -pub fn establish_connection(database_url: &str) -> SqliteConnection { - // dotenvy::dotenv().ok(); - // let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); - SqliteConnection::establish(&database_url) - .unwrap_or_else(|_| panic!("Error connecting to {}", database_url)) -} - pub fn populate_missing_london_date_ids( conn: &mut SqliteConnection, ) -> Result<(), RepositoryError> { diff --git a/src-tauri/src/download.rs b/src-tauri/src/download.rs index 5f0a3d2..3015499 100644 --- a/src-tauri/src/download.rs +++ b/src-tauri/src/download.rs @@ -1,12 +1,6 @@ -use std::{ - cmp, - error::Error, - future::Future, - sync::{Arc, Mutex}, -}; +use std::{cmp, error::Error, future::Future, sync::Arc}; use chrono::{Duration, Local, NaiveDate, NaiveDateTime}; -use diesel::SqliteConnection; use log::{debug, error, info}; use serde::Serialize; use tauri::{async_runtime, AppHandle}; @@ -25,6 +19,7 @@ use crate::{ }, RepositoryError, }, + db::SqliteConnectionPool, utils::{emit_event, get_or_create_energy_profile}, AppError, AppState, }; @@ -55,7 +50,7 @@ where T: EnergyDataProvider, { data_provider: Arc, - connection: Arc>, + connection_pool: SqliteConnectionPool, } impl DataLoader for ElectricityConsumptionDataLoader @@ -84,7 +79,8 @@ where fn insert_data(&self, data: Vec) -> Result<(), Self::InsertError> { if data.len() > 0 { - SqliteElectricityConsumptionRepository::new(self.connection.clone()).insert(data)?; + SqliteElectricityConsumptionRepository::new(self.connection_pool.clone()) + .insert(data)?; } Ok(()) @@ -97,7 +93,7 @@ where T: EnergyDataProvider, { data_provider: Arc, - connection: Arc>, + connection_pool: SqliteConnectionPool, } impl DataLoader for ElectricityTariffDataLoader @@ -134,7 +130,7 @@ where }) .collect(); - SqliteElectricityTariffRepository::new(self.connection.clone()).insert(tps)?; + SqliteElectricityTariffRepository::new(self.connection_pool.clone()).insert(tps)?; Ok(()) } @@ -146,7 +142,7 @@ where T: EnergyDataProvider, { data_provider: Arc, - connection: Arc>, + connection_pool: SqliteConnectionPool, } impl DataLoader for GasConsumptionDataLoader @@ -171,7 +167,7 @@ where fn insert_data(&self, data: Vec) -> Result<(), Self::InsertError> { if data.len() > 0 { - SqliteGasConsumptionRepository::new(self.connection.clone()).insert(data)?; + SqliteGasConsumptionRepository::new(self.connection_pool.clone()).insert(data)?; } Ok(()) @@ -184,7 +180,7 @@ where T: EnergyDataProvider, { data_provider: Arc, - connection: Arc>, + connection_pool: SqliteConnectionPool, } impl DataLoader for GasTariffDataLoader @@ -221,7 +217,7 @@ where }) .collect(); - SqliteGasTariffRepository::new(self.connection.clone()).insert(tps)?; + SqliteGasTariffRepository::new(self.connection_pool.clone()).insert(tps)?; Ok(()) } @@ -380,18 +376,18 @@ where let electricity_consumption_data_loader = ElectricityConsumptionDataLoader { data_provider: data_provider.clone(), - connection: app_state.db.clone(), + connection_pool: app_state.db_pool.clone(), }; let electricity_tariff_data_loader = ElectricityTariffDataLoader { data_provider: data_provider.clone(), - connection: app_state.db.clone(), + connection_pool: app_state.db_pool.clone(), }; let app_handle_clone = app_handle.clone(); check_for_new_data( - app_state.db.clone(), + app_state.db_pool.clone(), "electricity", "kWh", |until_date_time| async move { @@ -418,18 +414,18 @@ where let gas_consumption_data_loader = GasConsumptionDataLoader { data_provider: data_provider.clone(), - connection: app_state.db.clone(), + connection_pool: app_state.db_pool.clone(), }; let gas_tariff_data_loader = GasTariffDataLoader { data_provider: data_provider.clone(), - connection: app_state.db.clone(), + connection_pool: app_state.db_pool.clone(), }; let app_handle_clone = app_handle.clone(); check_for_new_data( - app_state.db.clone(), + app_state.db_pool.clone(), "gas", "kWh", |until_date_time| async move { @@ -458,7 +454,7 @@ where } async fn check_for_new_data( - connection: Arc>, + connection_pool: SqliteConnectionPool, profile_name: &str, base_unit: &str, download_action: F, @@ -467,7 +463,7 @@ where F: FnOnce(NaiveDateTime) -> Fut, Fut: Future>, { - let profile = get_or_create_energy_profile(connection.clone(), profile_name, base_unit)?; + let profile = get_or_create_energy_profile(connection_pool.clone(), profile_name, base_unit)?; if !profile.is_active { info!( @@ -481,7 +477,7 @@ where let last_date_retrieved = download_action(until_date_time).await?; - let repository = SqliteEnergyProfileRepository::new(connection); + let repository = SqliteEnergyProfileRepository::new(connection_pool); repository .update_energy_profile( diff --git a/src-tauri/src/main.rs b/src-tauri/src/main.rs index e7d1103..e6396db 100644 --- a/src-tauri/src/main.rs +++ b/src-tauri/src/main.rs @@ -3,6 +3,7 @@ use app_settings::{AppSettings, SETTINGS_FILE}; use clients::glowmarkt::GlowmarktDataProviderError; +use diesel::r2d2::{ConnectionManager, Pool}; use diesel::SqliteConnection; use log::{debug, error}; use std::env; @@ -22,7 +23,7 @@ use commands::glowmarkt::*; use commands::mqtt::*; use commands::profiles::*; -use crate::db::populate_missing_london_date_ids; +use crate::db::{populate_missing_london_date_ids, SqliteConnectionPool}; use crate::mqtt::start_mqtt_listener; use crate::utils::MqttSettings; use crate::utils::{get_mqtt_settings_opt, MqttAppSettings}; @@ -40,7 +41,7 @@ mod serde_utils; mod utils; struct AppState { - db: Arc>, + db_pool: SqliteConnectionPool, downloading: Arc>, client_available: Arc>, app_settings: Arc>, @@ -51,7 +52,7 @@ struct AppState { impl Clone for AppState { fn clone(&self) -> Self { Self { - db: self.db.clone(), + db_pool: self.db_pool.clone(), downloading: self.downloading.clone(), client_available: self.client_available.clone(), app_settings: self.app_settings.clone(), @@ -140,10 +141,23 @@ fn main() { let db_path = app_data_dir.join("db.sqlite"); - let mut connection = - db::establish_connection(db_path.to_str().expect("db path needed")); - db::run_migrations(&mut connection); - populate_missing_london_date_ids(&mut connection)?; + let connection_manager = ConnectionManager::::new( + db_path.to_str().expect("db path needed"), + ); + + let db_connection_pool = Pool::builder() + .max_size(10) + .build(connection_manager) + .expect("Failed to create database connection pool"); + + { + let mut connection = db_connection_pool + .get() + .expect("Failed to get connection from pool"); + + db::run_migrations(&mut connection); + populate_missing_london_date_ids(&mut connection)?; + } let store = app.store(SETTINGS_FILE)?; @@ -159,7 +173,7 @@ fn main() { let (tx, rx) = tokio::sync::mpsc::channel::(1); let app_state = AppState { - db: Arc::new(Mutex::new(connection)), + db_pool: db_connection_pool, downloading: Arc::new(Mutex::new(false)), client_available: Arc::new(Mutex::new(false)), app_settings: Arc::new(Mutex::new(app_settings)), diff --git a/src-tauri/src/utils.rs b/src-tauri/src/utils.rs index 0bef73c..0b26cee 100644 --- a/src-tauri/src/utils.rs +++ b/src-tauri/src/utils.rs @@ -1,8 +1,5 @@ -use std::sync::{Arc, Mutex}; - use chrono::{Datelike, NaiveDate, NaiveDateTime, TimeZone, Utc}; use chrono_tz::Europe::London; -use diesel::SqliteConnection; use keyring_core::Entry; use serde::{Deserialize, Serialize}; use tauri::{AppHandle, Emitter, Manager}; @@ -12,6 +9,7 @@ use crate::{ clients::glowmarkt::GlowmarktDataProvider, commands::{ApiError, APP_SERVICE_NAME}, data::energy_profile::{EnergyProfile, EnergyProfileRepository, SqliteEnergyProfileRepository}, + db::SqliteConnectionPool, AppError, AppState, MqttMessage, }; @@ -63,11 +61,11 @@ where } pub fn get_or_create_energy_profile( - connection: Arc>, + connection_pool: SqliteConnectionPool, name: &str, base_unit: &str, ) -> Result { - let repository = SqliteEnergyProfileRepository::new(connection); + let repository = SqliteEnergyProfileRepository::new(connection_pool); repository.get_energy_profile(name).or_else(|get_error| { repository