diff --git a/src/context.rs b/src/context.rs index ce849a0..506dd13 100644 --- a/src/context.rs +++ b/src/context.rs @@ -3,7 +3,6 @@ use crate::config::Config; use crate::polling::git::GitFetcher; use crate::repository::SqliteRepository; -use sqlx::SqlitePool; use std::sync::Arc; use tokio_util::sync::CancellationToken; @@ -16,9 +15,6 @@ pub struct SharedContext { /// Repository for data access. pub repository: Arc, - /// SQLx connection pool. - pub db_pool: SqlitePool, - /// Token to signal task cancellation. pub token: CancellationToken, diff --git a/src/handler.rs b/src/handler.rs index 5222955..8859e2d 100644 --- a/src/handler.rs +++ b/src/handler.rs @@ -69,7 +69,6 @@ mod tests { let state = AppState { config: std::sync::Arc::new(config), repository: std::sync::Arc::new(crate::repository::SqliteRepository::new(pool.clone())), - db_pool: pool.clone(), }; let payload = CreateSubscription { source_repo_url: RepoUrl::new("https://github.com/org/repo".to_string()).unwrap(), @@ -147,7 +146,6 @@ mod tests { let state = AppState { config: std::sync::Arc::new(config), repository: std::sync::Arc::new(crate::repository::SqliteRepository::new(pool.clone())), - db_pool: pool.clone(), }; // Try getting a non-existent subscription @@ -176,7 +174,6 @@ mod tests { let state = AppState { config: std::sync::Arc::new(config), repository: std::sync::Arc::new(crate::repository::SqliteRepository::new(pool.clone())), - db_pool: pool.clone(), }; // Create 3 subscriptions @@ -234,7 +231,6 @@ mod tests { let state = AppState { config: std::sync::Arc::new(config), repository: std::sync::Arc::new(crate::repository::SqliteRepository::new(pool.clone())), - db_pool: pool.clone(), }; let payload = CreateSubscription { source_repo_url: RepoUrl::new("https://github.com/org/repo".to_string()).unwrap(), diff --git a/src/lib.rs b/src/lib.rs index 535eb29..642c590 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -10,7 +10,6 @@ )] use std::fs; -use std::str::FromStr; use std::time::Duration; use axum::{ @@ -26,7 +25,6 @@ use reqwest::Client; use rovo::Router as RovoRouter; use rovo::aide::openapi::OpenApi; use rovo::rovo; -use sqlx::sqlite::SqliteConnectOptions; use subtle::ConstantTimeEq; use tokio::signal; use tokio_util::sync::CancellationToken; @@ -74,16 +72,11 @@ type EngineTask = (Box, &'static str); /// Runs the server, delegating errors to the caller. pub async fn run_app(tracker: &TaskTracker, token: &CancellationToken) -> Result<(), FatalError> { let config = Config::load()?; - let pool = init_database(&config).await?; - let repository = std::sync::Arc::new(crate::repository::SqliteRepository::new(pool.clone())); + let repository = + std::sync::Arc::new(crate::repository::SqliteRepository::connect(&config.database).await?); let http_client = build_http_client(&config)?; - let ctx = init_context( - repository.clone(), - pool.clone(), - config.clone(), - token.clone(), - )?; + let ctx = init_context(repository.clone(), config.clone(), token.clone())?; crate::trigger::recover_stuck_tasks(&repository, &config) .await @@ -94,7 +87,7 @@ pub async fn run_app(tracker: &TaskTracker, token: &CancellationToken) -> Result crate::engine::start_engine(engine, message, tracker); } - let app = build_router(repository, pool, &config); + let app = build_router(repository, &config); run_server(app, &ctx.config, token.clone()).await?; @@ -120,27 +113,9 @@ pub fn log_dotenv_status(loaded: bool) { ); } -/// Initializes the database pool. -async fn init_database(config: &Config) -> Result { - let options = SqliteConnectOptions::from_str(config.database.url.as_str())? - .foreign_keys(true) - .journal_mode(sqlx::sqlite::SqliteJournalMode::Wal); - - let pool = sqlx::sqlite::SqlitePoolOptions::new() - .acquire_timeout(config.database.timeout) - .connect_with(options) - .await?; - - // Ensures database schema is up to date in all environments. - sqlx::migrate!().run(&pool).await?; - - Ok(pool) -} - /// Initializes the shared application context. fn init_context( repository: std::sync::Arc, - pool: sqlx::SqlitePool, config: Config, token: CancellationToken, ) -> Result { @@ -156,7 +131,6 @@ fn init_context( Ok(SharedContext { config, repository, - db_pool: pool, token, git_fetcher: std::sync::Arc::new(git_fetcher), }) @@ -333,13 +307,11 @@ impl OnResponse for HttpRequestOnResponse { /// Builds the application router. pub fn build_router( repository: std::sync::Arc, - pool: sqlx::SqlitePool, config: &Config, ) -> Router { let state = AppState { config: std::sync::Arc::new(config.clone()), repository, - db_pool: pool, }; let mut api = OpenApi::default(); diff --git a/src/polling/mod.rs b/src/polling/mod.rs index 2234b84..1038ed0 100644 --- a/src/polling/mod.rs +++ b/src/polling/mod.rs @@ -261,7 +261,6 @@ mod tests { let ctx = SharedContext { config: crate::test_utils::create_test_config(), repository: std::sync::Arc::new(crate::repository::SqliteRepository::new(pool.clone())), - db_pool: pool.clone(), git_fetcher: mock_fetcher, token: CancellationToken::new(), }; @@ -314,7 +313,6 @@ mod tests { let ctx = SharedContext { config: crate::test_utils::create_test_config(), repository: std::sync::Arc::new(crate::repository::SqliteRepository::new(pool.clone())), - db_pool: pool.clone(), git_fetcher: Arc::new(crate::test_utils::MockGitFetcher { hash: CommitHash::new("b".repeat(40)).unwrap(), }), @@ -332,7 +330,6 @@ mod tests { let ctx = SharedContext { config: ctx.config, repository: std::sync::Arc::new(crate::repository::SqliteRepository::new(pool.clone())), - db_pool: pool.clone(), git_fetcher: mock_fetcher, token: ctx.token, }; @@ -402,7 +399,6 @@ mod tests { let ctx = SharedContext { config: crate::test_utils::create_test_config(), repository: std::sync::Arc::new(crate::repository::SqliteRepository::new(pool.clone())), - db_pool: pool.clone(), git_fetcher: std::sync::Arc::new(crate::test_utils::MockGitFetcher { hash: CommitHash::new("c".repeat(40)).unwrap(), }), diff --git a/src/repository/sqlite.rs b/src/repository/sqlite.rs index 965892b..5948bd8 100644 --- a/src/repository/sqlite.rs +++ b/src/repository/sqlite.rs @@ -1,6 +1,10 @@ //! SQLite implementation of the repository. +use std::str::FromStr; + +use crate::config::DatabaseConfig; use crate::domain::{BranchName, EventType, RepoUrl, TargetRepo}; +use crate::error::FatalError; use crate::model::{ Branch, CreateSubscription, Subscription, SubscriptionWithBranch, TriggerQueueItem, UpdateSubscription, @@ -13,6 +17,7 @@ use crate::repository::{ }; use async_trait::async_trait; use futures::future::BoxFuture; +use sqlx::sqlite::{SqliteConnectOptions, SqliteJournalMode, SqlitePoolOptions}; use sqlx::{SqliteConnection, SqlitePool}; #[derive(Debug)] @@ -23,6 +28,23 @@ pub struct SqliteRepository { } impl SqliteRepository { + /// Connects to the database described by `config`. + pub async fn connect(config: &DatabaseConfig) -> Result { + let options = SqliteConnectOptions::from_str(config.url.as_str())? + .foreign_keys(true) + .journal_mode(SqliteJournalMode::Wal); + + let pool = SqlitePoolOptions::new() + .acquire_timeout(config.timeout) + .connect_with(options) + .await?; + + // Ensures database schema is up to date in all environments. + sqlx::migrate!().run(&pool).await?; + + Ok(Self { pool }) + } + /// Creates a new [`SqliteRepository`] from a [`SqlitePool`]. pub fn new(pool: SqlitePool) -> Self { Self { pool } diff --git a/src/state.rs b/src/state.rs index 42aa824..7c6d821 100644 --- a/src/state.rs +++ b/src/state.rs @@ -14,7 +14,4 @@ pub struct AppState { /// Repository for data access. pub repository: Arc, - - /// SQLx connection pool for the SQLite database. - pub db_pool: sqlx::SqlitePool, } diff --git a/src/tests/api_routes.rs b/src/tests/api_routes.rs index 1dfce90..6b7feb5 100644 --- a/src/tests/api_routes.rs +++ b/src/tests/api_routes.rs @@ -13,7 +13,7 @@ async fn test_subscription_api_routes() { config.auth.allow_unauthenticated = true; let repository = Arc::new(SqliteRepository::new(pool.clone())); - let app = build_router(repository, pool, &config); + let app = build_router(repository, &config); // Test List Subscriptions (Empty) let response = app diff --git a/src/tests/auth_tests.rs b/src/tests/auth_tests.rs index 043233d..fa3fea9 100644 --- a/src/tests/auth_tests.rs +++ b/src/tests/auth_tests.rs @@ -16,7 +16,7 @@ async fn test_auth_no_key_configured_fails() { config.auth.allow_unauthenticated = false; let repository = Arc::new(SqliteRepository::new(pool.clone())); - let app = build_router(repository, pool, &config); + let app = build_router(repository, &config); let response = app .oneshot( @@ -40,7 +40,7 @@ async fn test_auth_allowed_unauthenticated_success() { config.auth.allow_unauthenticated = true; let repository = Arc::new(SqliteRepository::new(pool.clone())); - let app = build_router(repository, pool, &config); + let app = build_router(repository, &config); let response = app .oneshot( @@ -63,7 +63,7 @@ async fn test_auth_key_configured_success() { config.auth.api_key = Some(NonEmptyString::new("secret".to_string()).unwrap()); let repository = Arc::new(SqliteRepository::new(pool.clone())); - let app = build_router(repository, pool, &config); + let app = build_router(repository, &config); let response = app .oneshot( @@ -87,7 +87,7 @@ async fn test_auth_key_configured_mismatch() { config.auth.api_key = Some(NonEmptyString::new("secret".to_string()).unwrap()); let repository = Arc::new(SqliteRepository::new(pool.clone())); - let app = build_router(repository, pool, &config); + let app = build_router(repository, &config); let response = app .oneshot( @@ -111,7 +111,7 @@ async fn test_auth_key_configured_missing() { config.auth.api_key = Some(NonEmptyString::new("secret".to_string()).unwrap()); let repository = Arc::new(SqliteRepository::new(pool.clone())); - let app = build_router(repository, pool, &config); + let app = build_router(repository, &config); let response = app .oneshot( diff --git a/src/trigger/mod.rs b/src/trigger/mod.rs index 7ed3762..60a9f30 100644 --- a/src/trigger/mod.rs +++ b/src/trigger/mod.rs @@ -536,7 +536,6 @@ mod tests { repository: std::sync::Arc::new(crate::repository::SqliteRepository::new( pool.clone(), )), - db_pool: pool.clone(), token: CancellationToken::new(), git_fetcher: Arc::new(MockGitFetcher { hash: CommitHash::new("a".repeat(40)).unwrap(), @@ -623,7 +622,6 @@ mod tests { repository: std::sync::Arc::new(crate::repository::SqliteRepository::new( pool.clone(), )), - db_pool: pool.clone(), token: CancellationToken::new(), git_fetcher: Arc::new(MockGitFetcher { hash: CommitHash::new("a".repeat(40)).unwrap(),