diff --git a/Cargo.lock b/Cargo.lock index 90be89c1..16291452 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -62,6 +62,16 @@ dependencies = [ "rustversion", ] +[[package]] +name = "assert-json-diff" +version = "2.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47e4f2b81832e72834d7518d8487a0396a28cc408186a2e8854c0f98011faf12" +dependencies = [ + "serde", + "serde_json", +] + [[package]] name = "async-channel" version = "2.5.0" @@ -590,6 +600,24 @@ version = "2.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a4ae5f15dda3c708c0ade84bfee31ccab44a3da4f88015ed22f63732abe300c8" +[[package]] +name = "deadpool" +version = "0.12.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0be2b1d1d6ec8d846f05e137292d0b89133caf95ef33695424c09568bdd39b1b" +dependencies = [ + "deadpool-runtime", + "lazy_static", + "num_cpus", + "tokio", +] + +[[package]] +name = "deadpool-runtime" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "092966b41edc516079bdf31ec78a2e0588d1d0c08f78b91d8307215928642b2b" + [[package]] name = "debugid" version = "0.8.0" @@ -1020,6 +1048,12 @@ version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" +[[package]] +name = "hermit-abi" +version = "0.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e17592d60ebacc7d5e169f4663c5f84f9161cc90328abcfe8456f41e4dfcb284" + [[package]] name = "hex" version = "0.4.3" @@ -1679,6 +1713,16 @@ dependencies = [ "autocfg", ] +[[package]] +name = "num_cpus" +version = "1.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91df4bbde75afed763b708b7eee1e8e7651e02d97f6d5dd763e89367e957b23b" +dependencies = [ + "hermit-abi", + "libc", +] + [[package]] name = "objc2" version = "0.6.4" @@ -1995,10 +2039,11 @@ dependencies = [ [[package]] name = "opsqueue" -version = "0.41.0" +version = "0.42.0" dependencies = [ "anyhow", "arc-swap", + "async-trait", "axum", "axum-prometheus", "backon", @@ -2036,6 +2081,7 @@ dependencies = [ "tokio", "tokio-tungstenite 0.30.0", "tokio-util", + "tower", "tower-http 0.7.0", "tracing", "tracing-opentelemetry", @@ -2043,12 +2089,13 @@ dependencies = [ "url", "uuid", "ux", + "wiremock", "workspace-hack", ] [[package]] name = "opsqueue_python" -version = "0.41.0" +version = "0.42.0" dependencies = [ "anyhow", "chrono", @@ -4228,6 +4275,29 @@ dependencies = [ "memchr", ] +[[package]] +name = "wiremock" +version = "0.6.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08db1edfb05d9b3c1542e521aea074442088292f00b5f28e435c714a98f85031" +dependencies = [ + "assert-json-diff", + "base64", + "deadpool", + "futures", + "http", + "http-body-util", + "hyper", + "hyper-util", + "log", + "once_cell", + "regex", + "serde", + "serde_json", + "tokio", + "url", +] + [[package]] name = "wit-bindgen" version = "0.57.1" @@ -4262,6 +4332,7 @@ dependencies = [ "opentelemetry_sdk", "rand 0.10.2", "rand 0.9.5", + "regex", "regex-automata", "regex-syntax", "reqwest", diff --git a/Cargo.toml b/Cargo.toml index e8cca97d..ed8d964c 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -14,7 +14,7 @@ hakari-package = "workspace-hack" [workspace.package] edition = "2024" -version = "0.41.0" +version = "0.42.0" [workspace.dependencies] anyhow = { version = "1.0.102", default-features = false } diff --git a/opsqueue/Cargo.toml b/opsqueue/Cargo.toml index 75daf32e..944130ca 100644 --- a/opsqueue/Cargo.toml +++ b/opsqueue/Cargo.toml @@ -7,12 +7,12 @@ repository = "https://github.com/channable/opsqueue" license = "MIT" [lib] -name="opsqueue" -path="src/lib.rs" +name = "opsqueue" +path = "src/lib.rs" [[bin]] -name="opsqueue" -path="app/main.rs" +name = "opsqueue" +path = "app/main.rs" required-features = ["server-logic"] [dependencies] @@ -72,6 +72,7 @@ humantime.workspace = true dashmap.workspace = true sqlformat.workspace = true workspace-hack.workspace = true +async-trait = "0.1.91" # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html @@ -80,6 +81,8 @@ workspace = true [dev-dependencies] insta.workspace = true +wiremock = "0.6.5" +tower = "0.5.3" [[bench]] name = "chunks_select" @@ -98,8 +101,9 @@ server-logic = [ "dep:tower-http", "dep:axum-prometheus", "dep:sentry", - "dep:sentry-tracing" - ] + "dep:sentry-tracing", + "dep:reqwest", +] # Dependencies only in use by the client libraries: client-logic = [ "dep:reqwest", diff --git a/opsqueue/app/main.rs b/opsqueue/app/main.rs index 5baac612..8220988f 100644 --- a/opsqueue/app/main.rs +++ b/opsqueue/app/main.rs @@ -4,12 +4,15 @@ use opentelemetry_otlp::SpanExporter; use opentelemetry_resource_detectors::HostResourceDetector; use opentelemetry_resource_detectors::{OsResourceDetector, ProcessResourceDetector}; use opentelemetry_sdk::trace::{RandomIdGenerator, Sampler, SdkTracerProvider}; +use opsqueue::common::extension::{CoreApi, Extension}; +use opsqueue::delegation::extension::DelegationExtension; use opsqueue::tracing::as_dyn_error; use opsqueue::{common::submission::db::periodically_cleanup_old, config::Config, prometheus}; use std::{ sync::{Arc, atomic::AtomicBool}, time::Duration, }; +use tokio::sync::Notify; use tokio_util::sync::CancellationToken; use tracing::level_filters::LevelFilter; @@ -52,6 +55,24 @@ pub async fn async_main() { .await .expect("Timed out while initiating the database"); + let notify_on_insert = Arc::new(Notify::new()); + let (submission_status_changed_tx, _) = tokio::sync::broadcast::channel(128); + + let core_api = CoreApi::new( + db_pool.clone(), + notify_on_insert.clone(), + submission_status_changed_tx.clone(), + ); + + let mut extensions: Vec> = Vec::new(); + if let Some(delegation_server_url) = &config.delegation_server_url { + extensions.push(Box::new(DelegationExtension::new( + cancellation_token.clone(), + core_api, + delegation_server_url.clone(), + ))); + } + moro_local::async_scope!(|scope| { let checkpoint_handle = scope.spawn(db_pool.periodically_checkpoint_wal()); @@ -63,10 +84,14 @@ pub async fn async_main() { &cancellation_token, &app_healthy_flag, prometheus_config, + notify_on_insert, + submission_status_changed_tx, + &extensions, )); let max_age = config.max_submission_age.into(); - let cleanup_handle = scope.spawn(periodically_cleanup_old(db_pool.writer_pool(), max_age)); + + let cleanup_handle = scope.spawn(periodically_cleanup_old(&db_pool, max_age, &extensions)); let prometheus_handle = scope.spawn(prometheus::periodically_calculate_scaling_metrics( &db_pool, diff --git a/opsqueue/migrations/20260811150000_submissions_external_task.down.sql b/opsqueue/migrations/20260811150000_submissions_external_task.down.sql new file mode 100644 index 00000000..fc864983 --- /dev/null +++ b/opsqueue/migrations/20260811150000_submissions_external_task.down.sql @@ -0,0 +1 @@ +DROP TABLE submissions_external_task; diff --git a/opsqueue/migrations/20260811150000_submissions_external_task.up.sql b/opsqueue/migrations/20260811150000_submissions_external_task.up.sql new file mode 100644 index 00000000..677c0852 --- /dev/null +++ b/opsqueue/migrations/20260811150000_submissions_external_task.up.sql @@ -0,0 +1,6 @@ +CREATE TABLE submissions_external_task +( + submission_id BIGINT NOT NULL UNIQUE, + task_id TEXT NOT NULL UNIQUE, + last_status_sent TEXT +); diff --git a/opsqueue/opsqueue_example_database_schema.db b/opsqueue/opsqueue_example_database_schema.db index 9fbf6023..fe8c98cd 100644 Binary files a/opsqueue/opsqueue_example_database_schema.db and b/opsqueue/opsqueue_example_database_schema.db differ diff --git a/opsqueue/src/common/chunk.rs b/opsqueue/src/common/chunk.rs index acf2a7d0..5c7f756f 100644 --- a/opsqueue/src/common/chunk.rs +++ b/opsqueue/src/common/chunk.rs @@ -312,17 +312,21 @@ pub mod db { chunk_id: ChunkId, output_content: Option>, mut conn: impl WriterConnection, + submission_status_changed: &tokio::sync::broadcast::Sender, ) -> Result<(), E> { - let chunk_moved = conn + let (chunk_moved, completed_submission) = conn .transaction(move |mut tx| { Box::pin(async move { let chunk_moved = complete_chunk_raw(chunk_id, output_content, &mut tx).await?; + + let mut completed_submission = false; if chunk_moved { - crate::common::submission::db::maybe_complete_submission( - chunk_id.submission_id, - &mut tx, - ) - .await?; + completed_submission = + crate::common::submission::db::maybe_complete_submission( + chunk_id.submission_id, + &mut tx, + ) + .await?; } else { tracing::warn!( "Could not complete chunk {:?} because it was either: \ @@ -331,7 +335,10 @@ pub mod db { ); } - Result::>::Ok(chunk_moved) + Result::<(bool, bool), E>::Ok(( + chunk_moved, + completed_submission, + )) }) }) .await?; @@ -339,6 +346,9 @@ pub mod db { if chunk_moved { counter!(crate::prometheus::CHUNKS_COMPLETED_COUNTER).increment(1); } + if completed_submission { + let _ = submission_status_changed.send(chunk_id.submission_id); + } Ok(()) } @@ -413,6 +423,7 @@ pub mod db { failure: String, mut conn: impl WriterConnection, max_retries: u32, + submission_status_changed: &tokio::sync::broadcast::Sender, ) -> sqlx::Result { let failed_permanently = conn .transaction(move |mut tx| { @@ -461,6 +472,11 @@ pub mod db { }) }) .await?; + + if failed_permanently { + let _ = submission_status_changed.send(chunk_id.submission_id); + } + Ok(failed_permanently) } @@ -737,6 +753,33 @@ pub mod db { Ok(()) } + /// Delete all chunks belonging to the given submission. + /// + /// # Errors + /// + /// Returns an error if deletion failed. + #[tracing::instrument(skip(conn))] + pub async fn delete_chunks( + submission_id: SubmissionId, + mut conn: impl WriterConnection, + ) -> sqlx::Result<()> { + sqlx::query!( + " + DELETE FROM chunks WHERE chunks.submission_id = $1; + DELETE FROM chunks_paused WHERE chunks_paused.submission_id = $2; + DELETE FROM chunks_completed WHERE chunks_completed.submission_id = $3; + DELETE FROM chunks_failed WHERE chunks_failed.submission_id = $4; + ", + submission_id, + submission_id, + submission_id, + submission_id, + ) + .execute(conn.get_inner()) + .await?; + Ok(()) + } + /// Count chunks currently in progress. /// /// # Errors @@ -848,6 +891,55 @@ pub mod test { assert_eq!(count_chunks(&mut conn).await.unwrap(), 1); } + #[sqlx::test(migrator = "crate::MIGRATOR")] + pub async fn test_delete_chunks_removes_all_states(db: sqlx::SqlitePool) { + let db = WriterPool::new(db); + let mut conn = db.writer_conn().await.unwrap(); + let deleted_submission = SubmissionId::new(); + let retained_submission = SubmissionId::new(); + + for submission_id in [deleted_submission, retained_submission] { + sqlx::query!( + " + INSERT INTO chunks (submission_id, chunk_index) VALUES ($1, 0); + INSERT INTO chunks_paused (submission_id, chunk_index) VALUES ($2, 1); + INSERT INTO chunks_completed (submission_id, chunk_index, completed_at) + VALUES ($3, 2, julianday('now')); + INSERT INTO chunks_failed (submission_id, chunk_index, failed_at) + VALUES ($4, 3, julianday('now')); + ", + submission_id, + submission_id, + submission_id, + submission_id, + ) + .execute(conn.get_inner()) + .await + .unwrap(); + } + + delete_chunks(deleted_submission, &mut conn).await.unwrap(); + + let remaining = sqlx::query!( + r#" + SELECT submission_id AS "submission_id: SubmissionId" FROM chunks + UNION ALL SELECT submission_id FROM chunks_paused + UNION ALL SELECT submission_id FROM chunks_completed + UNION ALL SELECT submission_id FROM chunks_failed + "# + ) + .fetch_all(conn.get_inner()) + .await + .unwrap(); + assert_eq!( + remaining + .iter() + .map(|row| row.submission_id) + .collect::>(), + vec![retained_submission; 4] + ); + } + #[sqlx::test(migrator = "crate::MIGRATOR")] pub async fn test_get_chunk(db: sqlx::SqlitePool) { let db = WriterPool::new(db); @@ -956,6 +1048,8 @@ pub mod test { ) { let db = WriterPool::new(db); let mut conn = db.writer_conn().await.unwrap(); + let (submission_status_changed_tx, _) = tokio::sync::broadcast::channel(10); + let (submission, chunks) = Submission::from_vec( vec![Some("foo".into()), Some("bar".into()), Some("baz".into())], None, @@ -972,11 +1066,11 @@ pub mod test { .await .expect("insertion failed"); - let res = complete_chunk(chunk_id, None, &mut conn).await; - assert_matches!(res, Ok(())); + let res = complete_chunk(chunk_id, None, &mut conn, &submission_status_changed_tx).await; + assert!(res.is_ok()); - let res = complete_chunk(chunk_id, None, &mut conn).await; - assert_matches!(res, Ok(())); + let res = complete_chunk(chunk_id, None, &mut conn, &submission_status_changed_tx).await; + assert!(res.is_ok()); } #[sqlx::test(migrator = "crate::MIGRATOR")] @@ -1011,6 +1105,8 @@ pub mod test { ) { let db = WriterPool::new(db); let mut conn = db.writer_conn().await.unwrap(); + let (submission_status_changed_tx, _) = tokio::sync::broadcast::channel(10); + let (submission, chunks) = Submission::from_vec( vec![Some("foo".into()), Some("bar".into()), Some("baz".into())], None, @@ -1029,13 +1125,34 @@ pub mod test { let max_retries = 2; - let res = retry_or_fail_chunk(chunk_id, "kapot".into(), &mut conn, max_retries).await; + let res = retry_or_fail_chunk( + chunk_id, + "kapot".into(), + &mut conn, + max_retries, + &submission_status_changed_tx, + ) + .await; assert_matches!(res, Ok(false)); // Retry limit not yet reached. - let res = retry_or_fail_chunk(chunk_id, "kapot".into(), &mut conn, max_retries).await; + let res = retry_or_fail_chunk( + chunk_id, + "kapot".into(), + &mut conn, + max_retries, + &submission_status_changed_tx, + ) + .await; assert_matches!(res, Ok(true)); // Retry limit reached, submission is now permanently failed. - let res = retry_or_fail_chunk(chunk_id, "kapot".into(), &mut conn, max_retries).await; + let res = retry_or_fail_chunk( + chunk_id, + "kapot".into(), + &mut conn, + max_retries, + &submission_status_changed_tx, + ) + .await; assert_matches!(res, Ok(false)); // Submission was already failed, check that we ignore. } } diff --git a/opsqueue/src/common/extension.rs b/opsqueue/src/common/extension.rs new file mode 100644 index 00000000..a5e48edd --- /dev/null +++ b/opsqueue/src/common/extension.rs @@ -0,0 +1,109 @@ +use crate::E; +use crate::common::errors::{DatabaseError, E, SubmissionNotCancellable, SubmissionNotFound}; +use crate::common::submission::db::{cancel_submission, submission_status, unpause_submission}; +use crate::common::submission::{SubmissionId, SubmissionStatus}; +use crate::db; +use crate::db::{Connection, DBPools, WriterConnection}; +use std::sync::Arc; + +use async_trait::async_trait; +use axum::Router; +use tokio::sync::Notify; + +#[async_trait] +pub trait Extension: Send + Sync { + /// Returns whether the extension references the given submission. + /// Used to e.g. check if old submissions can be deleted safely. + async fn references_submission( + &self, + submission: SubmissionId, + conn: &mut db::conn::Writer, + ) -> sqlx::Result; + + /// Allows the Extension to register HTTP handlers. + /// Defaults to not registering anything. + fn bind_router(&self, router: Router) -> Router { + router + } +} + +#[derive(Debug, Clone)] +pub struct CoreApi { + pub pool: DBPools, + notify_on_insert: Arc, + submission_status_changed_tx: tokio::sync::broadcast::Sender, +} + +impl CoreApi { + #[must_use] + pub fn new( + pool: DBPools, + notify_on_insert: Arc, + submission_status_changed_tx: tokio::sync::broadcast::Sender, + ) -> Self { + Self { + pool, + notify_on_insert, + submission_status_changed_tx, + } + } + + /// Subscribe before starting a task that processes submission status changes. + #[must_use] + pub fn subscribe_submission_status_changes( + &self, + ) -> tokio::sync::broadcast::Receiver { + self.submission_status_changed_tx.subscribe() + } + + pub(crate) fn notify_submission_status_changed(&self, id: SubmissionId) { + let _ = self.submission_status_changed_tx.send(id); + } + + /// Unpauses the given submission. + /// + /// # Errors + /// + /// Returns a database error if the update fails, or `SubmissionNotFound` if the + /// submission is not paused. + pub async fn unpause_submission( + &self, + id: SubmissionId, + conn: impl WriterConnection, + ) -> Result<(), E> { + unpause_submission( + id, + conn, + &self.notify_on_insert, + &self.submission_status_changed_tx, + ) + .await + } + + /// Cancels the given submission. + /// + /// # Errors + /// + /// Returns a database error if the update fails, `SubmissionNotFound` if the + /// submission is missing, or `SubmissionNotCancellable` if it cannot be cancelled. + pub async fn cancel_submission( + &self, + id: SubmissionId, + conn: impl WriterConnection, + ) -> Result<(), E![DatabaseError, SubmissionNotFound, SubmissionNotCancellable]> { + cancel_submission(id, conn, &self.submission_status_changed_tx).await + } + + /// Gets the current status for the given submission. + /// + /// # Errors + /// + /// Returns a database error if the lookup fails. + pub async fn submission_status( + &self, + id: SubmissionId, + mut conn: impl Connection, + ) -> Result, DatabaseError> { + submission_status(id, &mut conn).await + } +} diff --git a/opsqueue/src/common/mod.rs b/opsqueue/src/common/mod.rs index 14fe5c5a..14e9ef45 100644 --- a/opsqueue/src/common/mod.rs +++ b/opsqueue/src/common/mod.rs @@ -4,6 +4,8 @@ use std::num::NonZero; pub mod chunk; pub mod errors; +#[cfg(feature = "server-logic")] +pub mod extension; pub mod submission; /// As values, we support the largest number value `SQLite` supports by itself, diff --git a/opsqueue/src/common/submission.rs b/opsqueue/src/common/submission.rs index a304c30a..1eab25fc 100644 --- a/opsqueue/src/common/submission.rs +++ b/opsqueue/src/common/submission.rs @@ -303,6 +303,14 @@ impl Submission { #[cfg(feature = "server-logic")] pub mod db { + use super::{ + Chunk, ChunkCount, ChunkIndex, DateTime, Duration, E, InitialSubmissionStatus, Metadata, + Submission, SubmissionCancelled, SubmissionCompleted, SubmissionFailed, SubmissionId, + SubmissionStatus, Utc, chunk, + }; + use crate::common::chunk::db::delete_chunks; + use crate::common::extension::Extension; + use crate::db::DBPools; use crate::tracing::as_dyn_error; use crate::{ common::{ @@ -313,17 +321,14 @@ pub mod db { }, submission::SubmissionPaused, }, - db::{Connection, True, WriterConnection, WriterPool}, + db::{Connection, True, WriterConnection}, }; use axum_prometheus::metrics::{counter, histogram}; use chunk::ChunkSize; + use futures::{StreamExt, TryStreamExt}; use sqlx::{Database, QueryBuilder, Sqlite, query, query_as, query_scalar}; - - use super::{ - Chunk, ChunkCount, ChunkIndex, DateTime, Duration, E, InitialSubmissionStatus, Metadata, - Submission, SubmissionCancelled, SubmissionCompleted, SubmissionFailed, SubmissionId, - SubmissionStatus, Utc, chunk, - }; + use std::sync::Arc; + use tokio::sync::Notify; impl<'q> sqlx::Encode<'q, Sqlite> for SubmissionId { fn encode_by_ref( @@ -542,6 +547,8 @@ pub mod db { pub async fn unpause_submission( id: SubmissionId, mut conn: impl WriterConnection, + notify_on_insert: &Arc, + submission_status_changed: &tokio::sync::broadcast::Sender, ) -> Result<(), E> { conn.transaction(move |mut tx| { Box::pin(async move { @@ -550,10 +557,16 @@ pub mod db { // NOTE: We need to check whether the submission is completed, because it might // be the case that we are unpausing a 0-chunk submission. maybe_complete_submission(id, &mut tx).await?; - Ok(()) + Ok::<(), E>(()) }) }) - .await + .await?; + + // Wake up any waiting consumers now that new chunks are available. + notify_on_insert.notify_waiters(); + let _ = submission_status_changed.send(id); + + Ok(()) } #[tracing::instrument(skip(conn))] @@ -1166,6 +1179,7 @@ pub mod db { pub async fn cancel_submission( id: SubmissionId, mut conn: impl WriterConnection, + submission_status_changed: &tokio::sync::broadcast::Sender, ) -> Result<(), E![DatabaseError, SubmissionNotFound, SubmissionNotCancellable]> { conn.transaction(move |mut tx| { Box::pin(async move { @@ -1173,10 +1187,13 @@ pub mod db { Ok(()) => Ok(()), Err(E::L(db_err)) => Err(E::L(db_err)), Err(E::R(not_found_err)) => { - // Submission was not found in the 'submissions' table, - // but it could still be in one of the other tables. + // Submission was not found in the 'submissions' or 'submissions_paused' + // tables, but it could still be in one of the other tables. match submission_status(id, &mut tx).await { Ok(None) => Err(E::R(E::L(not_found_err))), + Ok(Some(SubmissionStatus::Paused(submission))) => { + panic!("Failed to cancel paused submission {submission:?}") + } Ok(Some(SubmissionStatus::InProgress(submission))) => { panic!("Failed to cancel in progress submission {submission:?}") } @@ -1189,35 +1206,16 @@ pub mod db { Ok(Some(SubmissionStatus::Cancelled(submission))) => { Err(E::R(E::R(SubmissionNotCancellable::Cancelled(submission)))) } - Ok(Some(SubmissionStatus::Paused(_))) => { - // Paused submissions are cancellable. - cancel_paused_submission_notx(id, &mut tx).await.map_err( - |e| match e { - E::L(db_err) => E::L(db_err), - E::R(not_found) => E::R(E::L(not_found)), - }, - ) - } Err(db_err) => Err(E::L(db_err)), } } } }) }) - .await - } + .await?; + + let _ = submission_status_changed.send(id); - /// Do not call directly! Must be called inside a transaction. - /// - /// # Errors - /// - /// Returns an error if cancellation or chunk skipping fails. - async fn cancel_submission_notx( - id: SubmissionId, - mut conn: impl WriterConnection, - ) -> Result<(), E> { - cancel_submission_raw(id, &mut conn).await?; - super::chunk::db::skip_remaining_chunks(id, conn).await?; Ok(()) } @@ -1225,16 +1223,24 @@ pub mod db { /// /// # Errors /// - /// Returns [`DatabaseError`] if any SQL query fails. - /// - /// Returns [`SubmissionNotFound`] if the submission is not found in `submissions_paused`. - async fn cancel_paused_submission_notx( + /// Returns an error if cancellation or chunk skipping fails. + pub(crate) async fn cancel_submission_notx( id: SubmissionId, mut conn: impl WriterConnection, ) -> Result<(), E> { - cancel_paused_submission_raw(id, &mut conn).await?; - super::chunk::db::skip_remaining_paused_chunks(id, conn).await?; - Ok(()) + match cancel_submission_raw(id, &mut conn).await { + Ok(()) => { + super::chunk::db::skip_remaining_chunks(id, conn).await?; + Ok(()) + } + Err(E::R(_not_found)) => { + cancel_paused_submission_raw(id, &mut conn).await?; + super::chunk::db::skip_remaining_paused_chunks(id, conn).await?; + + Ok(()) + } + Err(E::L(db_err)) => Err(E::L(db_err)), + } } #[tracing::instrument(skip(conn))] @@ -1411,6 +1417,45 @@ pub mod db { Ok(()) } + /// Delete the given submission. + /// + /// # Errors + /// + /// Returns an error if deletion failed. + #[tracing::instrument(skip(conn))] + pub async fn delete_submission( + submission_id: SubmissionId, + mut conn: impl WriterConnection, + ) -> sqlx::Result<()> { + conn.transaction(move |mut tx| { + Box::pin(async move { + sqlx::query!( + " + DELETE FROM submissions_paused WHERE id = $1; + DELETE FROM submissions WHERE id = $2; + DELETE FROM submissions_cancelled WHERE id = $3; + DELETE FROM submissions_failed WHERE id = $4; + DELETE FROM submissions_completed WHERE id = $5; + DELETE FROM submissions_metadata WHERE submission_id = $6; + ", + submission_id, + submission_id, + submission_id, + submission_id, + submission_id, + submission_id, + ) + .execute(tx.get_inner()) + .await?; + + delete_chunks(submission_id, tx).await?; + + Ok(()) + }) + }) + .await + } + /// Count in-progress submissions. /// /// # Errors @@ -1504,95 +1549,61 @@ pub mod db { /// # Errors /// /// Returns an error if any cleanup statement in the transaction fails. - #[tracing::instrument(skip(conn))] + #[tracing::instrument(skip(db, extensions))] pub async fn cleanup_old( - mut conn: impl Connection, + db: &DBPools, older_than: DateTime, + extensions: &[Box], ) -> sqlx::Result<()> { tracing::info!("Cleaning up old completed/failed submissions..."); - conn.transaction(move |mut tx| { - Box::pin(async move { - // Clean up old submissions_metadata - query!( - "DELETE FROM submissions_metadata - WHERE submission_id IN ( - SELECT id FROM submissions_completed WHERE completed_at < julianday($1) - );", - older_than - ) - .execute(tx.get_inner()) - .await?; - query!( - "DELETE FROM submissions_metadata - WHERE submission_id IN ( - SELECT id FROM submissions_failed WHERE failed_at < julianday($1) - );", - older_than - ) - .execute(tx.get_inner()) - .await?; - query!( - "DELETE FROM submissions_metadata - WHERE submission_id IN ( - SELECT id FROM submissions_cancelled WHERE cancelled_at < julianday($1) - );", - older_than - ) - .execute(tx.get_inner()) - .await?; - // Clean up old submissions: - let n_submissions_completed = query!( - "DELETE FROM submissions_completed WHERE completed_at < julianday($1);", - older_than - ) - .execute(tx.get_inner()) - .await?.rows_affected(); - let n_submissions_failed = query!( - "DELETE FROM submissions_failed WHERE failed_at < julianday($1);", - older_than - ) - .execute(tx.get_inner()) - .await?.rows_affected(); - let n_submissions_cancelled = query!( - "DELETE FROM submissions_cancelled WHERE cancelled_at < julianday($1);", - older_than - ) - .execute(tx.get_inner()) - .await?.rows_affected(); + let mut read_conn = db.reader_conn().await?; + let mut old_submissions = query!( + r#" + SELECT id AS "id: SubmissionId" FROM submissions_completed WHERE completed_at < julianday($1) + UNION + SELECT id AS "id: SubmissionId" FROM submissions_failed WHERE failed_at < julianday($1) + UNION + SELECT id AS "id: SubmissionId" FROM submissions_cancelled WHERE cancelled_at < julianday($1) + "#, + older_than + ).fetch(read_conn.get_inner()).try_chunks(128); - let n_chunks_completed = query!( - "DELETE FROM chunks_completed WHERE completed_at < julianday($1);", - older_than - ) - .execute(tx.get_inner()) - .await?.rows_affected(); - let n_chunks_failed = query!( - "DELETE FROM chunks_failed WHERE failed_at < julianday($1);", - older_than - ) - .execute(tx.get_inner()) - .await?.rows_affected(); + let mut deleted_count: u64 = 0; - tracing::info!("Deleted {n_submissions_completed} completed submissions (with {n_chunks_completed} chunks completed)"); - tracing::info!("Deleted {n_submissions_failed} failed submissions (with {n_chunks_failed} chunks failed)"); - tracing::info!("Deleted {n_submissions_cancelled} cancelled submissions"); - Ok(()) - }) - }) - .await + while let Some(batch_res) = old_submissions.next().await { + let mut write_conn = db.writer_conn().await?; + + let batch = batch_res.map_err(|err| err.1)?; + + 'outer: for submission in batch { + for extension in extensions { + if extension + .references_submission(submission.id, &mut write_conn) + .await? + { + continue 'outer; + } + } + + delete_submission(submission.id, &mut write_conn).await?; + deleted_count += 1; + } + } + + tracing::debug!("Deleted {deleted_count} old submissions"); + Ok(()) } - pub async fn periodically_cleanup_old(db: &WriterPool, max_age: Duration) { + pub async fn periodically_cleanup_old( + db: &DBPools, + max_age: Duration, + extensions: &[Box], + ) { const PERIODIC_CLEANUP_INTERVAL: Duration = Duration::from_mins(1); loop { let cutoff = Utc::now() - max_age; - let res: sqlx::Result<()> = async move { - let mut conn = db.writer_conn().await?; - cleanup_old(&mut conn, cutoff).await?; - Ok(()) - } - .await; + let res = cleanup_old(db, cutoff, extensions).await; if let Err(e) = res { tracing::error!(error = as_dyn_error(&e), "Error during periodic cleanup"); } @@ -1604,16 +1615,17 @@ pub mod db { #[cfg(test)] #[cfg(feature = "server-logic")] pub mod test { + use crate::common::StrategicMetadataMap; + use crate::common::chunk::db::{count_chunks, count_chunks_failed, count_chunks_paused}; + use crate::db::{Connection as _, DBPools, WriterPool}; use chrono::Utc; use chunk::ChunkSize; use itertools::Itertools; use sqlformat::{FormatOptions, QueryParams, format}; use sqlx::{Execute, Row, Sqlite}; use std::assert_matches; - - use crate::common::StrategicMetadataMap; - use crate::common::chunk::db::{count_chunks, count_chunks_failed, count_chunks_paused}; - use crate::db::{Connection as _, WriterPool}; + use std::sync::Arc; + use tokio::sync::Notify; use super::db::*; use super::*; @@ -1867,141 +1879,164 @@ pub mod test { } #[sqlx::test(migrator = "crate::MIGRATOR")] - pub async fn test_cleanup_old(db: sqlx::SqlitePool) { - let db = WriterPool::new(db); - let mut conn = db.writer_conn().await.unwrap(); - - let chunks_contents = vec![Some("foo".into()), Some("bar".into()), Some("baz".into())]; - let old_one = insert_submission_from_chunks( - None, - chunks_contents.clone(), - None, - StrategicMetadataMap::default(), - ChunkSize::default(), - InitialSubmissionStatus::default(), - &mut conn, - ) - .await - .unwrap(); - let old_two = insert_submission_from_chunks( - None, - chunks_contents.clone(), - None, - StrategicMetadataMap::default(), - ChunkSize::default(), - InitialSubmissionStatus::default(), - &mut conn, - ) - .await - .unwrap(); - let old_three = insert_submission_from_chunks( - None, - chunks_contents.clone(), - None, - StrategicMetadataMap::default(), - ChunkSize::default(), - InitialSubmissionStatus::default(), - &mut conn, - ) - .await - .unwrap(); - let old_four_unfailed = insert_submission_from_chunks( - None, - chunks_contents.clone(), - None, - StrategicMetadataMap::default(), - ChunkSize::default(), - InitialSubmissionStatus::default(), - &mut conn, - ) - .await - .unwrap(); - - fail_submission(old_one, u63::new(0).into(), "Broken one".into(), &mut conn) + pub async fn test_cleanup_old( + pool_opts: sqlx::pool::PoolOptions, + conn_opts: sqlx::sqlite::SqliteConnectOptions, + ) { + let reader_pool = pool_opts + .clone() + .max_connections(16) + .connect_with(conn_opts.clone()) .await .unwrap(); - fail_submission(old_two, u63::new(0).into(), "Broken two".into(), &mut conn) + let writer_pool = pool_opts + .max_connections(1) + .connect_with(conn_opts) + .await + .unwrap(); + let db = DBPools::from_test_pools(&reader_pool, &writer_pool); + + let (cutoff_timestamp, old_four_unfailed) = { + let mut conn = db.writer_conn().await.unwrap(); + + let chunks_contents = vec![Some("foo".into()), Some("bar".into()), Some("baz".into())]; + let old_one = insert_submission_from_chunks( + None, + chunks_contents.clone(), + None, + StrategicMetadataMap::default(), + ChunkSize::default(), + InitialSubmissionStatus::default(), + &mut conn, + ) + .await + .unwrap(); + let old_two = insert_submission_from_chunks( + None, + chunks_contents.clone(), + None, + StrategicMetadataMap::default(), + ChunkSize::default(), + InitialSubmissionStatus::default(), + &mut conn, + ) + .await + .unwrap(); + let old_three = insert_submission_from_chunks( + None, + chunks_contents.clone(), + None, + StrategicMetadataMap::default(), + ChunkSize::default(), + InitialSubmissionStatus::default(), + &mut conn, + ) + .await + .unwrap(); + let old_four_unfailed = insert_submission_from_chunks( + None, + chunks_contents.clone(), + None, + StrategicMetadataMap::default(), + ChunkSize::default(), + InitialSubmissionStatus::default(), + &mut conn, + ) .await .unwrap(); - fail_submission( - old_three, - u63::new(0).into(), - "Broken three".into(), - &mut conn, - ) - .await - .unwrap(); - - // Ensure the clock is advanced ever so slightly. - // Not doing this makes the test flaky. - tokio::time::sleep(Duration::from_millis(1)).await; - - let cutoff_timestamp = Utc::now(); - - let too_new_one = insert_submission_from_chunks( - None, - chunks_contents.clone(), - None, - StrategicMetadataMap::default(), - ChunkSize::default(), - InitialSubmissionStatus::default(), - &mut conn, - ) - .await - .unwrap(); - let _too_new_two_unfailed = insert_submission_from_chunks( - None, - chunks_contents.clone(), - None, - StrategicMetadataMap::default(), - ChunkSize::default(), - InitialSubmissionStatus::default(), - &mut conn, - ) - .await - .unwrap(); - let too_new_three = insert_submission_from_chunks( - None, - chunks_contents.clone(), - None, - StrategicMetadataMap::default(), - ChunkSize::default(), - InitialSubmissionStatus::default(), - &mut conn, - ) - .await - .unwrap(); - fail_submission( - too_new_one, - u63::new(0).into(), - "Broken new one".into(), - &mut conn, - ) - .await - .unwrap(); - fail_submission( - too_new_three, - u63::new(0).into(), - "Broken new three".into(), - &mut conn, - ) - .await - .unwrap(); + fail_submission(old_one, u63::new(0).into(), "Broken one".into(), &mut conn) + .await + .unwrap(); + fail_submission(old_two, u63::new(0).into(), "Broken two".into(), &mut conn) + .await + .unwrap(); + fail_submission( + old_three, + u63::new(0).into(), + "Broken three".into(), + &mut conn, + ) + .await + .unwrap(); - assert_eq!(count_submissions_failed(&mut conn).await.unwrap(), 5); + // Ensure the clock is advanced ever so slightly. + // Not doing this makes the test flaky. + tokio::time::sleep(Duration::from_millis(1)).await; - let mut conn2 = db.writer_conn().await.unwrap(); - cleanup_old(&mut conn2, cutoff_timestamp).await.unwrap(); + let cutoff_timestamp = Utc::now(); - assert_eq!(count_submissions_failed(&mut conn).await.unwrap(), 2); + let too_new_one = insert_submission_from_chunks( + None, + chunks_contents.clone(), + None, + StrategicMetadataMap::default(), + ChunkSize::default(), + InitialSubmissionStatus::default(), + &mut conn, + ) + .await + .unwrap(); + let _too_new_two_unfailed = insert_submission_from_chunks( + None, + chunks_contents.clone(), + None, + StrategicMetadataMap::default(), + ChunkSize::default(), + InitialSubmissionStatus::default(), + &mut conn, + ) + .await + .unwrap(); + let too_new_three = insert_submission_from_chunks( + None, + chunks_contents.clone(), + None, + StrategicMetadataMap::default(), + ChunkSize::default(), + InitialSubmissionStatus::default(), + &mut conn, + ) + .await + .unwrap(); - let _sub1 = submission_status(old_four_unfailed, &mut conn) + fail_submission( + too_new_one, + u63::new(0).into(), + "Broken new one".into(), + &mut conn, + ) .await .unwrap(); - let _sub2 = submission_status(old_four_unfailed, &mut conn) + fail_submission( + too_new_three, + u63::new(0).into(), + "Broken new three".into(), + &mut conn, + ) .await .unwrap(); + + assert_eq!(count_submissions_failed(&mut conn).await.unwrap(), 5); + + (cutoff_timestamp, old_four_unfailed) + }; + + cleanup_old(&db, cutoff_timestamp, &[]).await.unwrap(); + + { + let mut conn = db.writer_conn().await.unwrap(); + + assert_eq!(count_submissions_failed(&mut conn).await.unwrap(), 2); + assert_eq!(count_chunks_failed(&mut conn).await.unwrap(), 6); + + let _sub1 = submission_status(old_four_unfailed, &mut conn) + .await + .unwrap(); + let _sub2 = submission_status(old_four_unfailed, &mut conn) + .await + .unwrap(); + } } #[sqlx::test(migrator = "crate::MIGRATOR")] @@ -2144,6 +2179,9 @@ pub mod test { pub async fn test_unpause_submission(db: sqlx::SqlitePool) { let db = WriterPool::new(db); let mut conn = db.writer_conn().await.unwrap(); + let (submission_status_changed_tx, _) = tokio::sync::broadcast::channel(10); + let notify_on_insert = Arc::new(Notify::new()); + let (submission, chunks) = Submission::from_vec( vec![Some("foo".into()), Some("bar".into()), Some("baz".into())], None, @@ -2159,7 +2197,14 @@ pub mod test { assert_eq!(count_chunks(&mut conn).await.unwrap(), 0); assert_eq!(count_chunks_paused(&mut conn).await.unwrap(), 3); - unpause_submission(submission.id, &mut conn).await.unwrap(); + unpause_submission( + submission.id, + &mut conn, + ¬ify_on_insert, + &submission_status_changed_tx, + ) + .await + .unwrap(); assert_eq!(count_submissions(&mut conn).await.unwrap(), 1); assert_eq!(count_submissions_paused(&mut conn).await.unwrap(), 0); assert_eq!(count_chunks(&mut conn).await.unwrap(), 3); @@ -2170,6 +2215,9 @@ pub mod test { pub async fn test_unpausing_a_zero_chunk_submission_completes_it(db: sqlx::SqlitePool) { let db = WriterPool::new(db); let mut conn = db.writer_conn().await.unwrap(); + let (submission_status_changed_tx, _) = tokio::sync::broadcast::channel(10); + let notify_on_insert = Arc::new(Notify::new()); + let (submission, chunks) = Submission::from_vec(vec![], None, ChunkSize::default()).unwrap(); insert_paused_submission(submission.clone(), chunks, &mut conn) @@ -2182,7 +2230,14 @@ pub mod test { assert_eq!(count_chunks(&mut conn).await.unwrap(), 0); assert_eq!(count_chunks_paused(&mut conn).await.unwrap(), 0); - unpause_submission(submission.id, &mut conn).await.unwrap(); + unpause_submission( + submission.id, + &mut conn, + ¬ify_on_insert, + &submission_status_changed_tx, + ) + .await + .unwrap(); assert_eq!(count_submissions(&mut conn).await.unwrap(), 0); assert_eq!(count_submissions_completed(&mut conn).await.unwrap(), 1); @@ -2195,6 +2250,8 @@ pub mod test { pub async fn test_cancel_paused_submission(db: sqlx::SqlitePool) { let db = WriterPool::new(db); let mut conn = db.writer_conn().await.unwrap(); + let (submission_status_changed_tx, _) = tokio::sync::broadcast::channel(10); + let (submission, chunks) = Submission::from_vec( vec![Some("foo".into()), Some("bar".into()), Some("baz".into())], None, @@ -2208,7 +2265,9 @@ pub mod test { assert_eq!(count_submissions_paused(&mut conn).await.unwrap(), 1); assert_eq!(count_chunks_paused(&mut conn).await.unwrap(), 3); - cancel_submission(submission.id, &mut conn).await.unwrap(); + cancel_submission(submission.id, &mut conn, &submission_status_changed_tx) + .await + .unwrap(); assert_eq!(count_submissions(&mut conn).await.unwrap(), 0); assert_eq!(count_submissions_paused(&mut conn).await.unwrap(), 0); diff --git a/opsqueue/src/config.rs b/opsqueue/src/config.rs index 26249cfb..80adaddb 100644 --- a/opsqueue/src/config.rs +++ b/opsqueue/src/config.rs @@ -105,6 +105,9 @@ pub struct Config { /// `lookup_submission_ids_by_strategic_metadata` request may return. #[arg(long, default_value_t = default_max_submissions_returned())] pub max_submissions_returned: MaxSubmissions, + + #[arg(long)] + pub delegation_server_url: Option, } impl Default for Config { @@ -122,6 +125,7 @@ impl Default for Config { let max_chunk_retries = 10; let max_submission_age = humantime::Duration::from_str("1 hour").expect("valid humantime"); let max_submissions_returned = default_max_submissions_returned(); + let delegation_server_url = None; Config { port, report_bound_port_pipe, @@ -133,6 +137,7 @@ impl Default for Config { max_chunk_retries, max_submission_age, max_submissions_returned, + delegation_server_url, } } } diff --git a/opsqueue/src/consumer/server/mod.rs b/opsqueue/src/consumer/server/mod.rs index 5cd9a224..62a56393 100644 --- a/opsqueue/src/consumer/server/mod.rs +++ b/opsqueue/src/consumer/server/mod.rs @@ -14,14 +14,14 @@ use axum_prometheus::metrics::{gauge, histogram}; use tokio::{select, sync::Notify}; use tokio_util::sync::CancellationToken; +use super::dispatcher::Dispatcher; +use crate::common::submission::SubmissionId; use crate::{ common::chunk::ChunkId, config::Config, db::{self, DBPools}, }; -use super::dispatcher::Dispatcher; - pub mod conn; pub mod state; @@ -37,10 +37,12 @@ pub async fn serve_for_tests( reservation_expiration: Duration, ) { let notify_on_insert = Arc::new(Notify::new()); + let submission_status_changed = tokio::sync::broadcast::channel(10).0; let config = Box::leak(Box::default()); let state = ServerState::new( pool, notify_on_insert, + submission_status_changed, cancellation_token.clone(), reservation_expiration, config, @@ -73,13 +75,18 @@ impl ServerState { pub fn new( pool: DBPools, notify_on_insert: Arc, + submission_status_changed: tokio::sync::broadcast::Sender, cancellation_token: CancellationToken, reservation_expiration: Duration, config: &'static Config, ) -> Self { let dispatcher = Dispatcher::new(reservation_expiration); - let (completer, completer_tx) = - Completer::new(pool.writer_pool(), &dispatcher, config.max_chunk_retries); + let (completer, completer_tx) = Completer::new( + pool.writer_pool(), + &dispatcher, + config.max_chunk_retries, + submission_status_changed, + ); Self { pool, completer: Some(completer), @@ -186,6 +193,7 @@ pub struct Completer { dispatcher: Dispatcher, count: usize, max_chunk_retries: u32, + submission_status_changed: tokio::sync::broadcast::Sender, } impl Completer { @@ -194,6 +202,7 @@ impl Completer { pool: &db::WriterPool, dispatcher: &Dispatcher, max_chunk_retries: u32, + submission_status_changed: tokio::sync::broadcast::Sender, ) -> (Self, tokio::sync::mpsc::Sender) { let (tx, rx) = tokio::sync::mpsc::channel(1024); let pool = pool.clone(); @@ -203,6 +212,7 @@ impl Completer { dispatcher: dispatcher.clone(), count: 0, max_chunk_retries, + submission_status_changed, }; (me, tx) } @@ -237,9 +247,13 @@ impl Completer { } => { // Even in the unlikely event that the DB write fails, // we still want to unreserve the chunk - let db_res = - crate::common::chunk::db::complete_chunk(id, output_content, &mut conn) - .await; + let db_res = crate::common::chunk::db::complete_chunk( + id, + output_content, + &mut conn, + &self.submission_status_changed, + ) + .await; reservations.lock().expect("No poison").remove(&id); if let Some(started_at) = self @@ -274,6 +288,7 @@ impl Completer { failure, &mut conn, self.max_chunk_retries, + &self.submission_status_changed, ) .await; reservations.lock().expect("No poison").remove(&id); diff --git a/opsqueue/src/db/mod.rs b/opsqueue/src/db/mod.rs index 4514b805..edc95746 100644 --- a/opsqueue/src/db/mod.rs +++ b/opsqueue/src/db/mod.rs @@ -200,6 +200,18 @@ impl DBPools { write_pool: Pool::new(pool.clone()), } } + + /// Create a `DBPools` instance from a single test pool. Only usable in tests. + #[cfg(test)] + pub(crate) fn from_test_pools( + read_pool: &sqlx::SqlitePool, + write_pool: &sqlx::SqlitePool, + ) -> Self { + DBPools { + read_pool: Pool::new(read_pool.clone()), + write_pool: Pool::new(write_pool.clone()), + } + } /// We check whether we can not only reach the DB but especially if we can run a transaction. /// /// This handles the case where for whatever reason some other thing holds the write lock for diff --git a/opsqueue/src/delegation/extension.rs b/opsqueue/src/delegation/extension.rs new file mode 100644 index 00000000..bf62f013 --- /dev/null +++ b/opsqueue/src/delegation/extension.rs @@ -0,0 +1,177 @@ +use crate::common::extension::{CoreApi, Extension}; +use crate::common::submission::SubmissionId; +use crate::db::Connection; +use crate::db::conn::{NoTransaction, Writer}; +use crate::delegation::server::ServerState; +use async_trait::async_trait; +use axum::Router; +use sqlx::query_scalar; +use tokio_util::sync::CancellationToken; + +pub struct DelegationExtension { + cancellation_token: CancellationToken, + core_api: CoreApi, + delegation_server_url: url::Url, +} + +impl DelegationExtension { + #[must_use] + pub fn new( + cancellation_token: CancellationToken, + core_api: CoreApi, + delegation_server_url: url::Url, + ) -> Self { + DelegationExtension { + cancellation_token, + core_api, + delegation_server_url, + } + } +} +#[async_trait] +impl Extension for DelegationExtension { + async fn references_submission( + &self, + submission: SubmissionId, + conn: &mut Writer, + ) -> sqlx::Result { + query_scalar!( + r#" + SELECT EXISTS(SELECT 1 FROM submissions_external_task WHERE submission_id = $1) AS "exists: bool" + "#, + submission + ) + .fetch_one(conn.get_inner()) + .await + } + + fn bind_router(&self, router: Router) -> Router { + let delegation_routes = ServerState::new( + self.cancellation_token.clone(), + self.core_api.clone(), + self.delegation_server_url.clone(), + ) + .run_background() + .build_router(); + + router.nest("/delegation", delegation_routes) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::common::StrategicMetadataMap; + use crate::common::chunk::ChunkSize; + use crate::common::submission::InitialSubmissionStatus; + use crate::common::submission::db::{ + count_submissions_completed, insert_submission_from_chunks, periodically_cleanup_old, + submission_status, + }; + use crate::db::DBPools; + use std::sync::Arc; + use std::time::Duration; + use tokio::sync::{Notify, broadcast}; + + #[sqlx::test(migrator = "crate::MIGRATOR")] + async fn periodic_cleanup_respects_external_task_references( + pool_opts: sqlx::pool::PoolOptions, + conn_opts: sqlx::sqlite::SqliteConnectOptions, + ) { + let reader_pool = pool_opts + .clone() + .max_connections(16) + .connect_with(conn_opts.clone()) + .await + .unwrap(); + let writer_pool = pool_opts + .max_connections(1) + .connect_with(conn_opts) + .await + .unwrap(); + let db = DBPools::from_test_pools(&reader_pool, &writer_pool); + + let (referenced, unreferenced) = { + let mut conn = db.writer_conn().await.unwrap(); + let referenced = insert_submission_from_chunks( + None, + vec![], + None, + StrategicMetadataMap::default(), + ChunkSize::default(), + InitialSubmissionStatus::InProgress, + &mut conn, + ) + .await + .unwrap(); + let unreferenced = insert_submission_from_chunks( + None, + vec![], + None, + StrategicMetadataMap::default(), + ChunkSize::default(), + InitialSubmissionStatus::InProgress, + &mut conn, + ) + .await + .unwrap(); + + sqlx::query!( + "INSERT INTO submissions_external_task (submission_id, task_id) VALUES ($1, 'task')", + referenced, + ) + .execute(conn.get_inner()) + .await + .unwrap(); + sqlx::query!( + "UPDATE submissions_completed SET completed_at = julianday('now', '-1 day')" + ) + .execute(conn.get_inner()) + .await + .unwrap(); + + (referenced, unreferenced) + }; + + let (sender, _) = broadcast::channel(1); + let extensions: Vec> = vec![Box::new(DelegationExtension::new( + CancellationToken::new(), + CoreApi::new(db.clone(), Arc::new(Notify::new()), sender), + "http://localhost/".parse().unwrap(), + ))]; + let cleanup_db = db.clone(); + let cleanup = tokio::spawn(async move { + periodically_cleanup_old(&cleanup_db, Duration::ZERO, &extensions).await; + }); + + let result = tokio::time::timeout(Duration::from_secs(5), async { + loop { + let mut conn = db.reader_conn().await.unwrap(); + let remaining = count_submissions_completed(&mut conn).await.unwrap(); + assert!(remaining > 0, "cleanup deleted the referenced submission"); + if remaining == 1 { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await; + cleanup.abort(); + assert!(cleanup.await.unwrap_err().is_cancelled()); + result.expect("periodic cleanup did not delete the unreferenced submission"); + + let mut conn = db.reader_conn().await.unwrap(); + assert!( + submission_status(referenced, &mut conn) + .await + .unwrap() + .is_some() + ); + assert!( + submission_status(unreferenced, &mut conn) + .await + .unwrap() + .is_none() + ); + } +} diff --git a/opsqueue/src/delegation/mod.rs b/opsqueue/src/delegation/mod.rs new file mode 100644 index 00000000..df860943 --- /dev/null +++ b/opsqueue/src/delegation/mod.rs @@ -0,0 +1,4 @@ +#[cfg(feature = "server-logic")] +pub mod extension; +#[cfg(feature = "server-logic")] +pub mod server; diff --git a/opsqueue/src/delegation/server.rs b/opsqueue/src/delegation/server.rs new file mode 100644 index 00000000..28c2e371 --- /dev/null +++ b/opsqueue/src/delegation/server.rs @@ -0,0 +1,1067 @@ +use crate::common::errors::{E, SubmissionNotFound}; +use crate::common::extension::CoreApi; +use crate::common::submission::{SubmissionId, SubmissionStatus}; +use crate::db::{Connection, DBPools, WriterConnection}; +use axum::extract::State; +use axum::http::StatusCode; +use axum::routing::post; +use axum::{Json, Router}; +use either::Either; +use futures::stream::BoxStream; +use itertools::Itertools; +use std::sync::Arc; +use tokio::select; +use tokio::sync::broadcast::error::{RecvError, TryRecvError}; +use tokio_util::sync::CancellationToken; +// use tower::ServiceExt; + +#[cfg(test)] +pub(crate) fn app_for_tests( + cancellation_token: &CancellationToken, + core_api: CoreApi, + delegation_server_url: url::Url, +) -> Router { + let router = ServerState::new(cancellation_token.clone(), core_api, delegation_server_url) + .run_background() + .build_router(); + + Router::new().nest("/delegation", router) +} + +#[derive(Debug, Clone)] +pub struct ServerState { + cancellation_token: CancellationToken, + core_api: CoreApi, + delegation_server_url: url::Url, + http_client: reqwest::Client, +} + +impl ServerState { + #[must_use] + pub fn new( + cancellation_token: CancellationToken, + core_api: CoreApi, + delegation_server_url: url::Url, + ) -> Self { + Self { + cancellation_token, + core_api, + delegation_server_url, + http_client: reqwest::Client::new(), + } + } + + #[must_use] + pub fn run_background(self) -> Self { + let submission_status_changed_rx = self.core_api.subscribe_submission_status_changes(); + let state = self.clone(); + tokio::spawn(run_in_background(state, submission_status_changed_rx)); + self + } + + pub fn build_router(self: ServerState) -> Router<()> { + Router::new() + .route("/submit", post(submit)) + .with_state(self) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, sqlx::Type)] +#[sqlx(type_name = "TEXT", rename_all = "snake_case")] +enum DelegatedJobStatus { + Paused, + InProgress, + Completed, + Failed, + Cancelled, +} + +#[derive(Debug, serde::Deserialize)] +#[serde(tag = "type", content = "contents")] +enum WorkerDelegationEvent { + #[serde(rename = "delegate")] + Delegate(Vec), + #[serde(rename = "kill")] + Kill(Vec), + #[serde(rename = "return")] + Return(Vec), +} +#[derive(Debug, serde::Serialize, serde::Deserialize)] +struct DelegatedJob { + task_id: String, + payload: DelegatedJobPayload, +} + +#[derive(Debug, serde::Serialize, serde::Deserialize)] +struct DelegatedJobPayload { + submission_id: SubmissionId, +} + +#[derive(Debug, serde::Serialize)] +#[serde(tag = "type", content = "contents")] +enum MasterDelegationEvent<'a> { + #[serde(rename = "updated")] + Updated(Vec>), + #[serde(rename = "completed")] + Completed(Vec>), +} + +#[derive(Debug, serde::Serialize)] +struct DelegatedJobUpdate<'a> { + task_id: &'a str, + status: DelegatedJobUpdateStatus, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize)] +#[serde(rename_all = "lowercase")] +enum DelegatedJobUpdateStatus { + Queued, + Running, +} + +#[derive(Debug, serde::Serialize)] +struct DelegatedJobCompletion<'a> { + task_id: &'a str, + completion: DelegatedJobCompletionStatus, +} + +#[derive(Debug, serde::Serialize)] +#[serde(tag = "status")] +enum DelegatedJobCompletionStatus { + #[serde(rename = "success")] + Success, + #[serde(rename = "failure")] + Failure { failure_reason: FailureReason }, +} + +#[derive(Debug, serde::Serialize)] +#[serde(rename_all = "lowercase")] +enum FailureReason { + Unknown, + Forced, +} + +#[tracing::instrument(level = "debug", skip(state))] +async fn submit( + State(state): State, + Json(events): Json>, +) -> Result { + let mut conn = state.core_api.pool.writer_conn().await.map_err(|e| { + tracing::error!("DB error acquiring writer connection: {e:?}"); + StatusCode::INTERNAL_SERVER_ERROR + })?; + + // Note that we don't need a transaction here. Say we fail half-way, JM will just retry. Since + // all operations in `handle_worker_events` are idempotent, that would be fine. + handle_worker_events(&state.core_api, &mut conn, events) + .await + .map_err(|e: sqlx::Error| { + tracing::error!("DB error handling events: {e:?}"); + StatusCode::INTERNAL_SERVER_ERROR + })?; + + Ok(StatusCode::ACCEPTED) +} + +async fn handle_worker_events( + core_api: &CoreApi, + conn: &mut impl WriterConnection, + events: Vec, +) -> sqlx::Result<()> { + for event in events { + match event { + WorkerDelegationEvent::Delegate(delegations) => { + for delegation in delegations { + handle_delegate_event(core_api, &mut *conn, &delegation) + .await + .map_err(|e| { + tracing::error!("Error handling delegate event: {e:?}"); + e + })?; + } + } + WorkerDelegationEvent::Kill(task_ids) => { + for task_id in task_ids { + handle_kill_event(core_api, &mut *conn, &task_id) + .await + .map_err(|e| { + tracing::error!( + "Error handling kill event for task_id={task_id}: {e:?}" + ); + e + })?; + } + } + WorkerDelegationEvent::Return(_task_ids) => { + tracing::info!( + "Received 'return' delegation event, which is not yet implemented; ignoring." + ); + } + } + } + + Ok(()) +} + +#[tracing::instrument(level = "debug", skip(conn))] +async fn handle_delegate_event( + core_api: &CoreApi, + conn: &mut impl WriterConnection, + job: &DelegatedJob, +) -> sqlx::Result<()> { + let task_id = &job.task_id; + let submission_id = job.payload.submission_id; + + let rows_affected = insert_external_task(&mut *conn, submission_id, task_id).await?; + + if rows_affected == 0 { + tracing::debug!(%submission_id, %task_id, "External task was already registered"); + } + + match core_api.unpause_submission(submission_id, &mut *conn).await { + Ok(()) => {} + Err(E::R(SubmissionNotFound(_))) => { + // Note that this path could be hit when doing resubmitting a task with the same + // submission id. + tracing::debug!(%submission_id, "Submission was not in paused state; either already active/completed/failed or already cleaned up."); + // Send a notification, s.t. the status is promptly reported to the external system. + core_api.notify_submission_status_changed(submission_id); + } + Err(E::L(db_err)) => { + tracing::error!(%submission_id, "DB error unpausing submission: {db_err:?}"); + return Err(db_err.0); + } + } + + Ok(()) +} + +#[tracing::instrument(level = "debug", skip(conn))] +async fn handle_kill_event( + core_api: &CoreApi, + conn: &mut impl WriterConnection, + task_id: &str, +) -> sqlx::Result<()> { + let submission_id = sqlx::query_scalar!( + r#"SELECT submission_id AS "submission_id: SubmissionId" + FROM submissions_external_task + WHERE task_id = $1"#, + task_id, + ) + .fetch_optional(conn.get_inner()) + .await?; + + let Some(submission_id) = submission_id else { + tracing::warn!(%task_id, "Kill event for unknown task_id; ignoring"); + return Ok(()); + }; + + match core_api.cancel_submission(submission_id, conn).await { + Ok(()) => {} + Err(E::L(db_err)) => { + tracing::error!(%submission_id, "DB error cancelling submission: {db_err:?}"); + return Err(db_err.0); + } + Err(E::R(E::L(SubmissionNotFound(_)))) => { + tracing::warn!(%submission_id, "Submission not found when attempting to cancel; already gone"); + } + Err(E::R(E::R(_not_cancelable))) => { + tracing::debug!(%submission_id, "Submission was already cancelled"); + } + } + + Ok(()) +} + +const DELEGATION_BACKGROUND_LOOP_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5); + +enum SubmissionStatusChangeBatchError { + Closed, + Lagged, +} +async fn submission_status_change_batch( + submission_status_changed: &mut tokio::sync::broadcast::Receiver, +) -> Result, SubmissionStatusChangeBatchError> { + match submission_status_changed.recv().await { + Ok(submission_id) => { + let mut submission_ids = vec![submission_id]; + loop { + match submission_status_changed.try_recv() { + Ok(submission_id) => submission_ids.push(submission_id), + Err(TryRecvError::Empty) => break, + Err(TryRecvError::Closed) => { + return Err(SubmissionStatusChangeBatchError::Closed); + } + Err(TryRecvError::Lagged(_)) => { + return Err(SubmissionStatusChangeBatchError::Lagged); + } + } + } + + Ok(submission_ids) + } + Err(RecvError::Closed) => Err(SubmissionStatusChangeBatchError::Closed), + Err(RecvError::Lagged(_)) => Err(SubmissionStatusChangeBatchError::Lagged), + } +} + +enum ReportTrigger { + FullScan, + Timeout, + Submissions(Vec), +} + +async fn run_in_background( + state: ServerState, + mut submission_status_changed_rx: tokio::sync::broadcast::Receiver, +) -> Result<(), ()> { + tracing::info!( + "Started delegation background loop. Updates will be sent to {}", + state.delegation_server_url + ); + + let mut trigger = ReportTrigger::FullScan; + + loop { + match report_submission_status(&state, trigger).await { + Ok(()) => {} + Err(e) => tracing::error!("Error in delegation background loop: {e:?}"), + } + + trigger = select! { + () = state.cancellation_token.cancelled() => break, + res = submission_status_change_batch(&mut submission_status_changed_rx) => { + match res { + Ok(submission_ids) => ReportTrigger::Submissions(submission_ids), + Err(SubmissionStatusChangeBatchError::Lagged) => ReportTrigger::FullScan, + Err(SubmissionStatusChangeBatchError::Closed) => break, + } + }, + () = tokio::time::sleep(DELEGATION_BACKGROUND_LOOP_TIMEOUT) => ReportTrigger::Timeout + }; + } + + Ok(()) +} + +async fn report_submission_status( + state: &ServerState, + trigger: ReportTrigger, +) -> Result<(), E> { + use futures::{StreamExt, TryStreamExt}; + let triggered_by_timeout = matches!(trigger, ReportTrigger::Timeout); + let update_specific_submissions = match trigger { + ReportTrigger::Submissions(submission_ids) => Some(submission_ids), + ReportTrigger::FullScan | ReportTrigger::Timeout => None, + }; + let mut conn_0 = state.core_api.pool.reader_conn().await.map_err(E::L)?; + let conn_1 = Arc::new(tokio::sync::Mutex::new( + state.core_api.pool.reader_conn().await.map_err(E::L)?, + )); + let core_api = Arc::new(state.core_api.clone()); + + let out_of_date_tasks = { + select_external_tasks(&mut conn_0, update_specific_submissions) + .try_filter_map({ + move |task| { + let conn_1 = conn_1.clone(); + let core_api = core_api.clone(); + async move { + let mut conn_1 = conn_1.lock().await; + let status = core_api + .submission_status(task.submission_id, &mut *conn_1) + .await + .map_err(|err| err.0)?; + Ok(out_of_date_task(task, (&status).into())) + } + } + }) + .try_chunks(128) + }; + + let mut task_count = 0; + + tokio::pin!(out_of_date_tasks); + + while let Some(batch_or_error) = out_of_date_tasks.next().await { + let batch = batch_or_error.map_err(|err| E::L(err.1))?; + + send_status_updates(state, batch.iter()) + .await + .map_err(E::R)?; + // Note that we don't need a transaction here. If `update_statuses_in_db` fails half-way, + // we will resend the update next iteration. Which is no big deal because + // `send_status_updates` is idempotent. + + update_statuses_in_db(&state.core_api.pool, batch.iter()) + .await + .map_err(E::L)?; + + task_count += batch.len(); + } + + if task_count > 0 && triggered_by_timeout { + tracing::warn!( + n_out_of_date_tasks = task_count, + "Delegation background loop was triggered by timeout with pending tasks; \ + possibly we are missing a submission_status_changed.send call" + ); + } + + Ok(()) +} + +async fn insert_external_task( + conn: &mut impl Connection, + submission_id: SubmissionId, + task_id: &str, +) -> sqlx::Result { + let rows_affected = sqlx::query!( + r#"INSERT INTO submissions_external_task (submission_id, task_id, last_status_sent) + SELECT $1 AS submission_id, $2 AS task_id, NULL AS last_status_sent + WHERE NOT EXISTS ( + SELECT TRUE + FROM submissions_external_task + WHERE submission_id = $1 AND task_id = $2 + )"#, + submission_id, + task_id, + ) + .execute(conn.get_inner()) + .await? + .rows_affected(); + + Ok(rows_affected) +} + +#[derive(Debug)] +struct ExternalTaskRow { + submission_id: SubmissionId, + task_id: String, + last_status_sent: Option, +} + +fn select_external_tasks( + conn: &mut impl Connection, + submission_ids: Option>, +) -> BoxStream<'_, sqlx::Result> { + if let Some(submission_ids) = submission_ids { + sqlx::query_as!( + ExternalTaskRow, + r#"SELECT + submission_id as "submission_id!: SubmissionId", + task_id, + last_status_sent as "last_status_sent: DelegatedJobStatus" + FROM submissions_external_task as t + WHERE t.submission_id IN (SELECT value FROM json_each($1)) + "#, + serde_json::to_string(&submission_ids).expect("Failed to serialize ids") + ) + .fetch(conn.get_inner()) + } else { + sqlx::query_as!( + ExternalTaskRow, + r#"SELECT + submission_id as "submission_id!: SubmissionId", + task_id, + last_status_sent as "last_status_sent: DelegatedJobStatus" + FROM submissions_external_task as t + "# + ) + .fetch(conn.get_inner()) + } +} + +struct OutOfDateTaskRow { + task_id: String, + current_status: DelegatedJobStatus, +} + +fn out_of_date_task( + task: ExternalTaskRow, + status: Option<&SubmissionStatus>, +) -> Option { + let current_status = match status { + Some(SubmissionStatus::Paused(_)) => DelegatedJobStatus::Paused, + Some(SubmissionStatus::InProgress(_)) => DelegatedJobStatus::InProgress, + Some(SubmissionStatus::Completed(_)) => DelegatedJobStatus::Completed, + Some(SubmissionStatus::Failed(_, _)) => DelegatedJobStatus::Failed, + Some(SubmissionStatus::Cancelled(_)) => DelegatedJobStatus::Cancelled, + None => { + let task_id = &task.task_id; + let submission_id = &task.submission_id; + // Log to Sentry, for visibility. + tracing::error!(%submission_id, %task_id, + "Got external task but could not find its submission. This could be a resubmitted \ + external task, where we already cleaned up the corresponding submission after \ + failure/completion. Assuming it completed successfully. If this delegated task \ + was manually resubmitted, consider regenerating the submission instead."); + DelegatedJobStatus::Completed + } + }; + (task.last_status_sent != Some(current_status)).then_some(OutOfDateTaskRow { + task_id: task.task_id, + current_status, + }) +} + +fn status_update( + task: &OutOfDateTaskRow, +) -> Either, DelegatedJobCompletion<'_>> { + match task.current_status { + DelegatedJobStatus::Paused => Either::Left(DelegatedJobUpdate { + task_id: &task.task_id, + status: DelegatedJobUpdateStatus::Queued, + }), + DelegatedJobStatus::InProgress => Either::Left(DelegatedJobUpdate { + task_id: &task.task_id, + status: DelegatedJobUpdateStatus::Running, + }), + DelegatedJobStatus::Completed => Either::Right(DelegatedJobCompletion { + task_id: &task.task_id, + completion: DelegatedJobCompletionStatus::Success, + }), + DelegatedJobStatus::Failed => Either::Right(DelegatedJobCompletion { + task_id: &task.task_id, + completion: DelegatedJobCompletionStatus::Failure { + failure_reason: FailureReason::Unknown, + }, + }), + DelegatedJobStatus::Cancelled => Either::Right(DelegatedJobCompletion { + task_id: &task.task_id, + completion: DelegatedJobCompletionStatus::Failure { + failure_reason: FailureReason::Forced, + }, + }), + } +} + +async fn update_statuses_in_db( + pool: &DBPools, + tasks: impl Iterator, +) -> sqlx::Result<()> { + let mut conn = pool.writer_conn().await?; + for task in tasks { + match task.current_status { + DelegatedJobStatus::Paused | DelegatedJobStatus::InProgress => { + update_last_status_sent(&mut conn, task).await?; + } + DelegatedJobStatus::Completed + | DelegatedJobStatus::Failed + | DelegatedJobStatus::Cancelled => { + delete_external_task(&mut conn, task).await?; + } + } + } + Ok(()) +} + +async fn update_last_status_sent( + mut conn: impl WriterConnection, + task: &OutOfDateTaskRow, +) -> sqlx::Result<()> { + sqlx::query!( + "UPDATE submissions_external_task SET last_status_sent = $1 WHERE task_id = $2", + task.current_status, + task.task_id, + ) + .execute(conn.get_inner()) + .await?; + + Ok(()) +} + +async fn delete_external_task( + mut conn: impl WriterConnection, + task: &OutOfDateTaskRow, +) -> sqlx::Result<()> { + sqlx::query!( + "DELETE FROM submissions_external_task WHERE task_id = $1", + task.task_id, + ) + .execute(conn.get_inner()) + .await?; + + Ok(()) +} + +async fn send_status_updates( + state: &ServerState, + events: impl Iterator, +) -> reqwest::Result<()> { + let (updates, completions): (Vec<_>, Vec<_>) = events.partition_map(status_update); + let events = vec![ + MasterDelegationEvent::Updated(updates), + MasterDelegationEvent::Completed(completions), + ]; + state + .http_client + .post( + state + .delegation_server_url + .join("/delegation/submit") + .unwrap(), + ) + .json(&events) + .send() + .await? + .error_for_status()?; + + Ok(()) +} + +#[cfg(test)] +#[cfg(feature = "server-logic")] +pub mod test { + use crate::common::StrategicMetadataMap; + use crate::common::chunk::db::{complete_chunk, retry_or_fail_chunk}; + use crate::common::chunk::{ChunkIndex, ChunkSize}; + use crate::common::extension::CoreApi; + use crate::common::submission::db::{ + count_submissions, count_submissions_cancelled, count_submissions_paused, + insert_submission_from_chunks, + }; + use crate::common::submission::{InitialSubmissionStatus, SubmissionId}; + use crate::db::{Connection, DBPools}; + use crate::delegation::server::{ + DelegatedJob, DelegatedJobPayload, app_for_tests, insert_external_task, + }; + use axum::http::Request; + use axum::{Router, body::Body}; + use http::{StatusCode, header}; + use serde_json::json; + use std::sync::{Arc, Mutex}; + use tokio::sync::{Notify, broadcast, oneshot}; + use tokio_util::sync::CancellationToken; + use tower::ServiceExt; + use wiremock::matchers::{body_partial_json, method, path}; + use wiremock::{Mock, MockServer, Respond, ResponseTemplate}; + + struct SignalResponder { + sender: Mutex>>, + response: ResponseTemplate, + } + + impl SignalResponder { + fn new(sender: oneshot::Sender<()>, response: ResponseTemplate) -> Self { + Self { + sender: Mutex::new(Some(sender)), + response, + } + } + } + + impl Respond for SignalResponder { + fn respond(&self, _request: &wiremock::Request) -> ResponseTemplate { + if let Ok(mut lock) = self.sender.lock() + && let Some(tx) = lock.take() + { + let _ = tx.send(()); + } + self.response.clone() + } + } + + async fn count_external_tasks(mut db: impl Connection) -> sqlx::Result { + let count = sqlx::query_scalar!("SELECT COUNT(*) as count FROM submissions_external_task;") + .fetch_one(db.get_inner()) + .await?; + Ok(u64::try_from(count).expect("COUNT(*) is always non-negative")) + } + + struct TestContext { + pool: DBPools, + external_server: MockServer, + core_api: CoreApi, + submission_status_changed: broadcast::Sender, + app: Router, + } + + impl TestContext { + async fn new( + pool_opts: sqlx::pool::PoolOptions, + conn_opts: sqlx::sqlite::SqliteConnectOptions, + ) -> Self { + let reader_pool = pool_opts + .clone() + .max_connections(16) + .connect_with(conn_opts.clone()) + .await + .unwrap(); + let writer_pool = pool_opts + .max_connections(1) + .connect_with(conn_opts) + .await + .unwrap(); + let pool = DBPools::from_test_pools(&reader_pool, &writer_pool); + let external_server = MockServer::start().await; + let (submission_status_changed, _) = broadcast::channel(128); + let core_api = CoreApi::new( + pool.clone(), + Arc::new(Notify::new()), + submission_status_changed.clone(), + ); + let cancellation_token = CancellationToken::new(); + let app = app_for_tests( + &cancellation_token, + core_api.clone(), + external_server.uri().parse().unwrap(), + ); + Self { + pool, + external_server, + core_api, + submission_status_changed, + app, + } + } + + async fn post_worker_events(&self, events: serde_json::Value) { + let request = Request::builder() + .uri("/delegation/submit") + .method("POST") + .header(header::CONTENT_TYPE, "application/json") + .body(Body::from(events.to_string())) + .unwrap(); + let response = self.app.clone().oneshot(request).await.unwrap(); + assert_eq!( + response.status(), + StatusCode::ACCEPTED, + "request failed: {response:?}" + ); + } + } + + async fn wait_for_update(rx: oneshot::Receiver<()>) { + tokio::time::timeout(std::time::Duration::from_secs(2), rx) + .await + .expect("Timed out waiting for HTTP request") + .expect("Sender dropped without signaling"); + } + + #[sqlx::test(migrator = "crate::MIGRATOR")] + pub async fn test_job_delegation( + pool_opts: sqlx::pool::PoolOptions, + conn_opts: sqlx::sqlite::SqliteConnectOptions, + ) { + let test = TestContext::new(pool_opts, conn_opts).await; + + let submission = { + let mut conn = test.pool.writer_conn().await.unwrap(); + + let chunks_contents = vec![Some("foo".into())]; + insert_submission_from_chunks( + None, + chunks_contents.clone(), + None, + StrategicMetadataMap::default(), + ChunkSize::default(), + InitialSubmissionStatus::Paused, + &mut conn, + ) + .await + .unwrap() + }; + + { + let mut conn = test.pool.reader_conn().await.unwrap(); + assert_eq!(count_submissions_paused(&mut conn).await.unwrap(), 1); + assert_eq!(count_external_tasks(&mut conn).await.unwrap(), 0); + } + + let (tx, rx) = oneshot::channel::<()>(); + Mock::given(method("POST")) + .and(path("/delegation/submit")) + .and(body_partial_json( + json!([{"type": "updated", "contents": [{"task_id": "test", "status": "running"}]}, {"type": "completed", "contents": []}]), + )) + .respond_with(SignalResponder::new(tx, ResponseTemplate::new(202))) + .expect(1) + .mount(&test.external_server) + .await; + + test.post_worker_events(json!([{"type": "delegate", "contents": [DelegatedJob { + task_id: "test".to_string(), + payload: DelegatedJobPayload { submission_id: submission }, + }]}])) + .await; + + wait_for_update(rx).await; + + { + let mut conn = test.pool.reader_conn().await.unwrap(); + assert_eq!(count_external_tasks(&mut conn).await.unwrap(), 1); + assert_eq!(count_submissions(&mut conn).await.unwrap(), 1); + } + } + + #[sqlx::test(migrator = "crate::MIGRATOR")] + pub async fn test_job_kill( + pool_opts: sqlx::pool::PoolOptions, + conn_opts: sqlx::sqlite::SqliteConnectOptions, + ) { + let test = TestContext::new(pool_opts, conn_opts).await; + + { + let mut conn = test.pool.writer_conn().await.unwrap(); + + let chunks_contents = vec![Some("foo".into())]; + let submission = insert_submission_from_chunks( + None, + chunks_contents.clone(), + None, + StrategicMetadataMap::default(), + ChunkSize::default(), + InitialSubmissionStatus::Paused, + &mut conn, + ) + .await + .unwrap(); + + insert_external_task(&mut conn, submission, "test") + .await + .unwrap(); + }; + + let (tx, rx) = oneshot::channel::<()>(); + Mock::given(method("POST")) + .and(path("/delegation/submit")) + .and(body_partial_json(json!([{"type": "updated", "contents": []}, {"type": "completed", "contents": [{"task_id": "test", "completion": {"status": "failure", "failure_reason": "forced"}}]}]))) + .respond_with(SignalResponder::new(tx, ResponseTemplate::new(202))) + .expect(1) + .mount(&test.external_server) + .await; + + test.post_worker_events(json!([{"type": "kill", "contents": ["test"]}])) + .await; + + { + let mut conn = test.pool.reader_conn().await.unwrap(); + assert_eq!(count_submissions_cancelled(&mut conn).await.unwrap(), 1); + assert_eq!(count_submissions(&mut conn).await.unwrap(), 0); + } + + wait_for_update(rx).await; + + // Wait for background loop to remove external tasks; + tokio::time::sleep(std::time::Duration::from_millis(200)).await; + + { + let mut conn = test.pool.reader_conn().await.unwrap(); + assert_eq!(count_external_tasks(&mut conn).await.unwrap(), 0); + } + } + + #[sqlx::test(migrator = "crate::MIGRATOR")] + pub async fn test_unpause_update( + pool_opts: sqlx::pool::PoolOptions, + conn_opts: sqlx::sqlite::SqliteConnectOptions, + ) { + let test = TestContext::new(pool_opts, conn_opts).await; + + let submission = { + let mut conn = test.pool.writer_conn().await.unwrap(); + + let chunks_contents = vec![Some("foo".into())]; + let submission = insert_submission_from_chunks( + None, + chunks_contents.clone(), + None, + StrategicMetadataMap::default(), + ChunkSize::default(), + InitialSubmissionStatus::Paused, + &mut conn, + ) + .await + .unwrap(); + + insert_external_task(&mut conn, submission, "test") + .await + .unwrap(); + + submission + }; + + let (tx, rx) = oneshot::channel::<()>(); + Mock::given(method("POST")) + .and(path("/delegation/submit")) + .and(body_partial_json( + json!([{"type": "updated", "contents": [{"task_id": "test", "status": "running"}]}, {"type": "completed", "contents": []}]), + )) + .respond_with(SignalResponder::new(tx, ResponseTemplate::new(202))) + .expect(1) + .mount(&test.external_server) + .await; + + { + let mut conn = test.pool.writer_conn().await.unwrap(); + test.core_api + .unpause_submission(submission, &mut conn) + .await + .unwrap(); + } + + wait_for_update(rx).await; + } + + #[sqlx::test(migrator = "crate::MIGRATOR")] + pub async fn test_complete_update( + pool_opts: sqlx::pool::PoolOptions, + conn_opts: sqlx::sqlite::SqliteConnectOptions, + ) { + let test = TestContext::new(pool_opts, conn_opts).await; + + let submission = { + let mut conn = test.pool.writer_conn().await.unwrap(); + + let chunks_contents = vec![Some("foo".into())]; + let submission = insert_submission_from_chunks( + None, + chunks_contents.clone(), + None, + StrategicMetadataMap::default(), + ChunkSize::default(), + InitialSubmissionStatus::InProgress, + &mut conn, + ) + .await + .unwrap(); + + insert_external_task(&mut conn, submission, "test") + .await + .unwrap(); + + submission + }; + + let (tx, rx) = oneshot::channel::<()>(); + Mock::given(method("POST")) + .and(path("/delegation/submit")) + .and(body_partial_json( + json!([{"type": "updated", "contents": []}, {"type": "completed", "contents": [{"task_id": "test", "completion": {"status": "success"}}]}]), + )) + .respond_with(SignalResponder::new(tx, ResponseTemplate::new(202))) + .expect(1) + .mount(&test.external_server) + .await; + + { + let mut conn = test.pool.writer_conn().await.unwrap(); + complete_chunk( + (submission, ChunkIndex::zero()).into(), + None, + &mut conn, + &test.submission_status_changed, + ) + .await + .unwrap(); + } + + wait_for_update(rx).await; + } + + #[sqlx::test(migrator = "crate::MIGRATOR")] + pub async fn test_fail_update( + pool_opts: sqlx::pool::PoolOptions, + conn_opts: sqlx::sqlite::SqliteConnectOptions, + ) { + let test = TestContext::new(pool_opts, conn_opts).await; + + let submission = { + let mut conn = test.pool.writer_conn().await.unwrap(); + + let chunks_contents = vec![Some("foo".into())]; + let submission = insert_submission_from_chunks( + None, + chunks_contents.clone(), + None, + StrategicMetadataMap::default(), + ChunkSize::default(), + InitialSubmissionStatus::InProgress, + &mut conn, + ) + .await + .unwrap(); + + insert_external_task(&mut conn, submission, "test") + .await + .unwrap(); + + submission + }; + + let (tx, rx) = oneshot::channel::<()>(); + Mock::given(method("POST")) + .and(path("/delegation/submit")) + .and(body_partial_json(json!([{"type": "updated", "contents": []}, {"type": "completed", "contents": [{"task_id": "test", "completion": {"status": "failure", "failure_reason": "unknown"}}]}]))) + .respond_with(SignalResponder::new(tx, ResponseTemplate::new(202))) + .expect(1) + .mount(&test.external_server) + .await; + + { + let mut conn = test.pool.writer_conn().await.unwrap(); + retry_or_fail_chunk( + (submission, ChunkIndex::zero()).into(), + "extreme error".to_owned(), + &mut conn, + 0, + &test.submission_status_changed, + ) + .await + .unwrap(); + } + + wait_for_update(rx).await; + } + + #[sqlx::test(migrator = "crate::MIGRATOR")] + pub async fn test_cancel_update( + pool_opts: sqlx::pool::PoolOptions, + conn_opts: sqlx::sqlite::SqliteConnectOptions, + ) { + let test = TestContext::new(pool_opts, conn_opts).await; + + let submission = { + let mut conn = test.pool.writer_conn().await.unwrap(); + + let chunks_contents = vec![Some("foo".into())]; + let submission = insert_submission_from_chunks( + None, + chunks_contents.clone(), + None, + StrategicMetadataMap::default(), + ChunkSize::default(), + InitialSubmissionStatus::InProgress, + &mut conn, + ) + .await + .unwrap(); + + insert_external_task(&mut conn, submission, "test") + .await + .unwrap(); + + submission + }; + + let (tx, rx) = oneshot::channel::<()>(); + Mock::given(method("POST")) + .and(path("/delegation/submit")) + .and(body_partial_json(json!([{"type": "updated", "contents": []}, {"type": "completed", "contents": [{"task_id": "test", "completion": {"status": "failure", "failure_reason": "forced"}}]}]))) + .respond_with(SignalResponder::new(tx, ResponseTemplate::new(202))) + .expect(1) + .mount(&test.external_server) + .await; + + { + let mut conn = test.pool.writer_conn().await.unwrap(); + test.core_api + .cancel_submission(submission, &mut conn) + .await + .unwrap(); + } + + wait_for_update(rx).await; + } +} diff --git a/opsqueue/src/lib.rs b/opsqueue/src/lib.rs index 7adfd4c8..eb555b58 100644 --- a/opsqueue/src/lib.rs +++ b/opsqueue/src/lib.rs @@ -38,6 +38,7 @@ pub mod prometheus; #[cfg(feature = "server-logic")] pub mod config; +pub mod delegation; /// The Opsqueue library's semantic version /// as written in the Rust packages's `Cargo.toml` diff --git a/opsqueue/src/producer/server.rs b/opsqueue/src/producer/server.rs index 2e51a3ad..9d0b9ae9 100644 --- a/opsqueue/src/producer/server.rs +++ b/opsqueue/src/producer/server.rs @@ -17,15 +17,21 @@ use super::common::{ChunkContents, InsertSubmission}; pub async fn serve_for_tests(database_pool: DBPools, server_addr: Box) { let max_submissions = crate::config::Config::default().max_submissions_returned; - ServerState::new(database_pool, Arc::new(Notify::new()), max_submissions) - .serve_for_tests(server_addr) - .await; + ServerState::new( + database_pool, + Arc::new(Notify::new()), + tokio::sync::broadcast::channel(10).0, + max_submissions, + ) + .serve_for_tests(server_addr) + .await; } #[derive(Debug, Clone)] pub struct ServerState { pool: DBPools, notify_on_insert: Arc, + submission_status_changed: tokio::sync::broadcast::Sender, max_submissions: MaxSubmissions, } @@ -33,11 +39,13 @@ impl ServerState { pub fn new( pool: DBPools, notify_on_insert: Arc, + submission_status_changed: tokio::sync::broadcast::Sender, max_submissions: MaxSubmissions, ) -> Self { ServerState { pool, notify_on_insert, + submission_status_changed, max_submissions, } } @@ -130,16 +138,15 @@ async fn cancel_submission( .writer_conn() .await .map_err(|e| ServerError(e.into()).into_response())?; - match submission::db::cancel_submission(submission_id, &mut conn).await { - Ok(()) => Ok(()), - Err(L(db_err)) => Err(ServerError(db_err.into()).into_response()), - Err(R(L(not_found_err))) => { - Err((StatusCode::NOT_FOUND, Json(not_found_err)).into_response()) - } - Err(R(R(not_cancellable_err))) => { - Err((StatusCode::CONFLICT, Json(not_cancellable_err)).into_response()) - } - } + submission::db::cancel_submission(submission_id, &mut conn, &state.submission_status_changed) + .await + .map_err(|err| match err { + L(db_err) => ServerError(db_err.into()).into_response(), + R(L(not_found_err)) => (StatusCode::NOT_FOUND, Json(not_found_err)).into_response(), + R(R(not_cancellable_err)) => { + (StatusCode::CONFLICT, Json(not_cancellable_err)).into_response() + } + }) } /// 200 if the submission was successfully unpaused. @@ -154,15 +161,17 @@ async fn unpause_submission( .writer_conn() .await .map_err(|e| ServerError(e.into()).into_response())?; - match submission::db::unpause_submission(submission_id, &mut conn).await { - Ok(()) => { - // Wake up any waiting consumers now that new chunks are available. - state.notify_on_insert.notify_waiters(); - Ok(()) - } - Err(L(db_err)) => Err(ServerError(db_err.into()).into_response()), - Err(R(not_found_err)) => Err((StatusCode::NOT_FOUND, Json(not_found_err)).into_response()), - } + submission::db::unpause_submission( + submission_id, + &mut conn, + &state.notify_on_insert, + &state.submission_status_changed, + ) + .await + .map_err(|err| match err { + L(db_err) => ServerError(db_err.into()).into_response(), + R(not_found_err) => (StatusCode::NOT_FOUND, Json(not_found_err)).into_response(), + }) } async fn submission_status( diff --git a/opsqueue/src/server.rs b/opsqueue/src/server.rs index 5f9ec428..408a7a33 100644 --- a/opsqueue/src/server.rs +++ b/opsqueue/src/server.rs @@ -1,5 +1,8 @@ //! Defines the HTTP endpoints that are used by both the `producer` and `consumer` APIs use crate::tracing::as_dyn_error; +use axum::{Router, routing::get}; +use backon::{BackoffBuilder, FibonacciBuilder}; +use http::{Response, StatusCode, header}; use std::{ any::Any, mem, @@ -7,13 +10,11 @@ use std::{ time::Duration, }; -use axum::{Router, routing::get}; -use backon::{BackoffBuilder, FibonacciBuilder}; -use http::{Response, StatusCode, header}; - +use crate::common::extension::Extension; +use crate::common::submission::SubmissionId; use crate::db::DBPools; use tokio::select; -use tokio::sync::Notify; +use tokio::sync::{Notify, broadcast}; use tokio_util::sync::CancellationToken; fn retry_policy() -> impl BackoffBuilder { @@ -25,6 +26,7 @@ fn retry_policy() -> impl BackoffBuilder { } #[cfg(feature = "server-logic")] +#[allow(clippy::too_many_arguments)] /// Start serving producer and consumer endpoints. /// /// # Errors @@ -38,18 +40,27 @@ pub async fn serve_producer_and_consumer( cancellation_token: &CancellationToken, app_healthy_flag: &Arc, prometheus_config: crate::prometheus::PrometheusConfig, + notify_on_insert: Arc, + submission_status_changed_tx: broadcast::Sender, + extensions: &[Box], ) -> Result<(), std::io::Error> { use backon::Retryable; - (|| async { - let router = build_router( - config, - pool.clone(), - reservation_expiration, - cancellation_token, - app_healthy_flag.clone(), - prometheus_config.clone(), - ); + let router = build_router( + config, + pool.clone(), + reservation_expiration, + cancellation_token, + app_healthy_flag.clone(), + prometheus_config, + notify_on_insert, + submission_status_changed_tx, + extensions, + ); + + (|| { + let router = router.clone(); + async move { let listener = tokio::net::TcpListener::bind(server_addr).await?; match listener.local_addr() { Ok(addr) => { @@ -73,7 +84,7 @@ pub async fn serve_producer_and_consumer( .with_graceful_shutdown(cancellation_token.clone().cancelled_owned()) .await?; Ok(()) - }) + }}) .retry(retry_policy()) .notify(|e, d| { tracing::error!( @@ -90,6 +101,7 @@ pub async fn serve_producer_and_consumer( } #[cfg(feature = "server-logic")] +#[allow(clippy::too_many_arguments)] pub fn build_router( config: &'static crate::config::Config, pool: DBPools, @@ -97,12 +109,14 @@ pub fn build_router( cancellation_token: &CancellationToken, app_healthy_flag: Arc, prometheus_config: crate::prometheus::PrometheusConfig, + notify_on_insert: Arc, + submission_status_changed_tx: broadcast::Sender, + extensions: &[Box], ) -> Router<()> { - let notify_on_insert = Arc::new(Notify::new()); - let consumer_routes = crate::consumer::server::ServerState::new( pool.clone(), notify_on_insert.clone(), + submission_status_changed_tx.clone(), cancellation_token.clone(), reservation_expiration, config, @@ -112,14 +126,19 @@ pub fn build_router( let producer_routes = crate::producer::server::ServerState::new( pool, notify_on_insert, + submission_status_changed_tx, config.max_submissions_returned, ) .build_router(); - let routes = Router::new() + let mut routes = Router::new() .nest("/producer", producer_routes) .nest("/consumer", consumer_routes); + for extension in extensions { + routes = extension.bind_router(routes); + } + let tracing_middleware = tower_http::trace::TraceLayer::new_for_http() .make_span_with(|request: &http::Request<_>| { use tracing_opentelemetry::OpenTelemetrySpanExt; diff --git a/workspace-hack/Cargo.toml b/workspace-hack/Cargo.toml index e647bbb6..a1b41c07 100644 --- a/workspace-hack/Cargo.toml +++ b/workspace-hack/Cargo.toml @@ -26,7 +26,8 @@ futures-channel = { version = "0.3", features = ["sink"] } futures-io = { version = "0.3" } futures-sink = { version = "0.3" } futures-util = { version = "0.3", features = ["channel", "io", "sink"] } -hyper = { version = "1", features = ["client", "http1", "http2", "server"] } +hyper = { version = "1", features = ["full"] } +hyper-util = { version = "0.1", features = ["client-legacy", "http1", "http2", "server", "service"] } libsqlite3-sys = { version = "0.30", default-features = false, features = ["bundled", "pkg-config", "unlock_notify", "vcpkg"] } log = { version = "0.4", default-features = false, features = ["std"] } num-traits = { version = "0.2", default-features = false, features = ["std"] } @@ -35,8 +36,9 @@ opentelemetry-http = { version = "0.32", default-features = false, features = [" opentelemetry_sdk = { version = "0.32", default-features = false, features = ["metrics", "rt-tokio", "trace"] } rand-274715c4dabd11b0 = { package = "rand", version = "0.9" } rand-93f6ce9d446188ac = { package = "rand", version = "0.10" } -regex-automata = { version = "0.4", default-features = false, features = ["dfa-build", "meta", "std", "unicode-perl", "unicode-word-boundary"] } -regex-syntax = { version = "0.8", default-features = false, features = ["std", "unicode-perl"] } +regex = { version = "1" } +regex-automata = { version = "0.4", default-features = false, features = ["dfa-build", "dfa-onepass", "hybrid", "meta", "nfa-backtrack", "perf-inline", "perf-literal", "std", "unicode"] } +regex-syntax = { version = "0.8" } reqwest = { version = "0.13", default-features = false, features = ["blocking", "http2", "json", "rustls", "stream"] } rustls-pki-types = { version = "1", features = ["std"] } serde = { version = "1", features = ["alloc", "derive", "rc"] } @@ -51,6 +53,7 @@ thiserror = { version = "2" } tokio = { version = "1", features = ["fs", "io-util", "macros", "net", "rt-multi-thread", "signal", "sync", "time"] } tokio-stream = { version = "0.1", features = ["fs"] } tokio-util = { version = "0.7", features = ["codec", "io", "rt", "time"] } +tower = { version = "0.5", default-features = false, features = ["balance", "buffer", "limit", "load-shed", "log"] } tracing-core = { version = "0.1" } typenum = { version = "1", default-features = false, features = ["const-generics"] } url = { version = "2", features = ["serde"] } @@ -89,9 +92,9 @@ url = { version = "2", features = ["serde"] } [target.x86_64-unknown-linux-gnu.dependencies] bitflags = { version = "2", default-features = false, features = ["std"] } -hyper-util = { version = "0.1", features = ["client-legacy", "client-proxy", "http1", "http2", "server", "service"] } +hyper-util = { version = "0.1", default-features = false, features = ["client-proxy"] } libc = { version = "0.2", features = ["extra_traits"] } -tower = { version = "0.5", default-features = false, features = ["balance", "buffer", "limit", "load-shed", "log", "retry", "timeout"] } +tower = { version = "0.5", default-features = false, features = ["retry", "timeout"] } tower-http = { version = "0.6", features = ["follow-redirect"] } [target.x86_64-unknown-linux-gnu.build-dependencies]