From 64f0fe40f8d182c3085b2370938023064cc94181 Mon Sep 17 00:00:00 2001 From: Sem Mulder Date: Wed, 22 Jul 2026 15:02:33 +0200 Subject: [PATCH 01/11] Use u64 in count_* functions u63 forces us to wrap literals in `u63::new`, and we need to convert to u64 at actual usage sites anyway. --- opsqueue/src/common/chunk.rs | 46 +++++++++++-------- opsqueue/src/common/submission.rs | 75 +++++++++++++------------------ opsqueue/src/producer/client.rs | 7 ++- opsqueue/src/producer/server.rs | 4 +- opsqueue/src/prometheus.rs | 4 +- 5 files changed, 64 insertions(+), 72 deletions(-) diff --git a/opsqueue/src/common/chunk.rs b/opsqueue/src/common/chunk.rs index c27c17ec..a56515d5 100644 --- a/opsqueue/src/common/chunk.rs +++ b/opsqueue/src/common/chunk.rs @@ -226,7 +226,6 @@ impl Chunk { pub mod db { use super::{ Chunk, ChunkCompleted, ChunkFailed, ChunkId, ChunkIndex, ChunkSize, DateTime, SubmissionId, - Utc, u63, }; use crate::common::errors::{ChunkNotFound, DatabaseError, E, SubmissionNotFound}; use crate::db::{Connection, True, WriterConnection}; @@ -620,13 +619,16 @@ pub mod db { /// # Errors /// /// Returns an error if the count query fails. + /// + /// # Panics + /// + /// Panics if `COUNT(*)` returns a negative value, which `SQLite` never does. #[tracing::instrument(skip(db))] - pub async fn count_chunks(mut db: impl Connection) -> sqlx::Result { + pub async fn count_chunks(mut db: impl Connection) -> sqlx::Result { let count = sqlx::query_scalar!("SELECT COUNT(1) as count FROM chunks;") .fetch_one(db.get_inner()) .await?; - let count = u63::new(count.cast_unsigned()); - Ok(count) + Ok(u64::try_from(count).expect("COUNT(*) is always non-negative")) } /// Count completed chunks. @@ -634,13 +636,16 @@ pub mod db { /// # Errors /// /// Returns an error if the count query fails. + /// + /// # Panics + /// + /// Panics if `COUNT(*)` returns a negative value, which `SQLite` never does. #[tracing::instrument(skip(db))] - pub async fn count_chunks_completed(mut db: impl Connection) -> sqlx::Result { + pub async fn count_chunks_completed(mut db: impl Connection) -> sqlx::Result { let count = sqlx::query_scalar!("SELECT COUNT(1) as count FROM chunks_completed;") .fetch_one(db.get_inner()) .await?; - let count = u63::new(count.cast_unsigned()); - Ok(count) + Ok(u64::try_from(count).expect("COUNT(*) is always non-negative")) } /// Count failed chunks. @@ -648,13 +653,16 @@ pub mod db { /// # Errors /// /// Returns an error if the count query fails. + /// + /// # Panics + /// + /// Panics if `COUNT(*)` returns a negative value, which `SQLite` never does. #[tracing::instrument(skip(db))] - pub async fn count_chunks_failed(mut db: impl Connection) -> sqlx::Result { + pub async fn count_chunks_failed(mut db: impl Connection) -> sqlx::Result { let count = sqlx::query_scalar!("SELECT COUNT(1) as count FROM chunks_failed;") .fetch_one(db.get_inner()) .await?; - let count = u63::new(count.cast_unsigned()); - Ok(count) + Ok(u64::try_from(count).expect("COUNT(*) is always non-negative")) } /// Looks up the number of operations in the backlog. @@ -692,11 +700,11 @@ pub mod test { vec![1, 2, 3, 4, 5].into(), ); - assert_eq!(count_chunks(&mut conn).await.unwrap(), u63::new(0)); + assert_eq!(count_chunks(&mut conn).await.unwrap(), 0); insert_chunk(chunk.clone(), &mut conn) .await .expect("Insert chunk failed"); - assert_eq!(count_chunks(&mut conn).await.unwrap(), u63::new(1)); + assert_eq!(count_chunks(&mut conn).await.unwrap(), 1); } #[sqlx::test(migrator = "crate::MIGRATOR")] @@ -752,12 +760,12 @@ pub mod test { .await .expect("complete chunk failed"); - assert_eq!(count_chunks(&mut conn).await.unwrap(), u63::new(0)); + assert_eq!(count_chunks(&mut conn).await.unwrap(), 0); assert_eq!( count_chunks_completed(&mut conn).await.unwrap(), - u63::new(1) + 1 ); - assert_eq!(count_chunks_failed(&mut conn).await.unwrap(), u63::new(0)); + assert_eq!(count_chunks_failed(&mut conn).await.unwrap(), 0); } #[sqlx::test(migrator = "crate::MIGRATOR")] @@ -775,7 +783,7 @@ pub mod test { .await .unwrap(); - assert_eq!(count_chunks(&mut conn).await.unwrap(), u63::new(1)); + assert_eq!(count_chunks(&mut conn).await.unwrap(), 1); conn.transaction(move |mut tx| { Box::pin(async move { @@ -824,11 +832,11 @@ pub mod test { .await .expect("Succeed chunk failed"); - assert_eq!(count_chunks(&mut conn).await.unwrap(), u63::new(0)); + assert_eq!(count_chunks(&mut conn).await.unwrap(), 0); assert_eq!( count_chunks_completed(&mut conn).await.unwrap(), - u63::new(0) + 0 ); - assert_eq!(count_chunks_failed(&mut conn).await.unwrap(), u63::new(1)); + assert_eq!(count_chunks_failed(&mut conn).await.unwrap(), 1); } } diff --git a/opsqueue/src/common/submission.rs b/opsqueue/src/common/submission.rs index 991744dd..a742b381 100644 --- a/opsqueue/src/common/submission.rs +++ b/opsqueue/src/common/submission.rs @@ -290,7 +290,6 @@ pub mod db { use axum_prometheus::metrics::{counter, histogram}; use chunk::ChunkSize; use sqlx::{QueryBuilder, Sqlite, query, query_scalar}; - use ux::u63; use super::{ Chunk, ChunkCount, ChunkIndex, DateTime, Duration, E, Metadata, Submission, @@ -1048,12 +1047,16 @@ pub mod db { /// # Errors /// /// Returns an error if the count query fails. + /// + /// # Panics + /// + /// Panics if `COUNT(*)` returns a negative value, which `SQLite` never does. #[tracing::instrument(skip(db))] - pub async fn count_submissions(mut db: impl Connection) -> sqlx::Result { + pub async fn count_submissions(mut db: impl Connection) -> sqlx::Result { let count = sqlx::query_scalar!("SELECT COUNT(1) as count FROM submissions;") .fetch_one(db.get_inner()) .await?; - Ok(u63::new(count.cast_unsigned())) + Ok(u64::try_from(count).expect("COUNT(*) is always non-negative")) } /// Count completed submissions. @@ -1061,12 +1064,16 @@ pub mod db { /// # Errors /// /// Returns an error if the count query fails. + /// + /// # Panics + /// + /// Panics if `COUNT(*)` returns a negative value, which `SQLite` never does. #[tracing::instrument(skip(db))] - pub async fn count_submissions_completed(mut db: impl Connection) -> sqlx::Result { + pub async fn count_submissions_completed(mut db: impl Connection) -> sqlx::Result { let count = sqlx::query_scalar!("SELECT COUNT(1) as count FROM submissions_completed;") .fetch_one(db.get_inner()) .await?; - Ok(u63::new(count.cast_unsigned())) + Ok(u64::try_from(count).expect("COUNT(*) is always non-negative")) } /// Count failed submissions. @@ -1074,12 +1081,16 @@ pub mod db { /// # Errors /// /// Returns an error if the count query fails. + /// + /// # Panics + /// + /// Panics if `COUNT(*)` returns a negative value, which `SQLite` never does. #[tracing::instrument(skip(db))] - pub async fn count_submissions_failed(mut db: impl Connection) -> sqlx::Result { + pub async fn count_submissions_failed(mut db: impl Connection) -> sqlx::Result { let count = sqlx::query_scalar!("SELECT COUNT(1) as count FROM submissions_failed;") .fetch_one(db.get_inner()) .await?; - Ok(u63::new(count.cast_unsigned())) + Ok(u64::try_from(count).expect("COUNT(*) is always non-negative")) } /// Transactionally removes all completed/failed submissions, @@ -1381,7 +1392,7 @@ pub mod test { let db = WriterPool::new(db); let mut conn = db.writer_conn().await.unwrap(); - assert_eq!(count_submissions(&mut conn).await.unwrap(), u63::new(0)); + assert_eq!(count_submissions(&mut conn).await.unwrap(), 0); let (submission, chunks) = Submission::from_vec( vec![Some("foo".into()), Some("bar".into()), Some("baz".into())], @@ -1393,7 +1404,7 @@ pub mod test { .await .expect("insertion failed"); - assert_eq!(count_submissions(&mut conn).await.unwrap(), u63::new(1)); + assert_eq!(count_submissions(&mut conn).await.unwrap(), 1); } #[sqlx::test(migrator = "crate::MIGRATOR")] @@ -1466,15 +1477,9 @@ pub mod test { .await .unwrap(); - assert_eq!(count_submissions(&mut conn).await.unwrap(), u63::new(0)); - assert_eq!( - count_submissions_completed(&mut conn).await.unwrap(), - u63::new(1) - ); - assert_eq!( - count_submissions_failed(&mut conn).await.unwrap(), - u63::new(0) - ); + assert_eq!(count_submissions(&mut conn).await.unwrap(), 0); + assert_eq!(count_submissions_completed(&mut conn).await.unwrap(), 1); + assert_eq!(count_submissions_failed(&mut conn).await.unwrap(), 0); } #[sqlx::test(migrator = "crate::MIGRATOR")] @@ -1499,15 +1504,9 @@ pub mod test { ) .await .unwrap(); - assert_eq!(count_submissions(&mut conn).await.unwrap(), u63::new(0)); - assert_eq!( - count_submissions_completed(&mut conn).await.unwrap(), - u63::new(0) - ); - assert_eq!( - count_submissions_failed(&mut conn).await.unwrap(), - u63::new(1) - ); + assert_eq!(count_submissions(&mut conn).await.unwrap(), 0); + assert_eq!(count_submissions_completed(&mut conn).await.unwrap(), 0); + assert_eq!(count_submissions_failed(&mut conn).await.unwrap(), 1); } #[sqlx::test(migrator = "crate::MIGRATOR")] @@ -1626,18 +1625,12 @@ pub mod test { .await .unwrap(); - assert_eq!( - count_submissions_failed(&mut conn).await.unwrap(), - u63::new(5) - ); + assert_eq!(count_submissions_failed(&mut conn).await.unwrap(), 5); let mut conn2 = db.writer_conn().await.unwrap(); cleanup_old(&mut conn2, cutoff_timestamp).await.unwrap(); - assert_eq!( - count_submissions_failed(&mut conn).await.unwrap(), - u63::new(2) - ); + assert_eq!(count_submissions_failed(&mut conn).await.unwrap(), 2); let _sub1 = submission_status(old_four_unfailed, &mut conn) .await @@ -1668,15 +1661,9 @@ pub mod test { .await .expect("insertion failed"); - assert_eq!(count_submissions(&mut conn).await.unwrap(), u63::new(0)); - assert_eq!( - count_submissions_completed(&mut conn).await.unwrap(), - u63::new(1) - ); - assert_eq!( - count_submissions_failed(&mut conn).await.unwrap(), - u63::new(0) - ); + assert_eq!(count_submissions(&mut conn).await.unwrap(), 0); + assert_eq!(count_submissions_completed(&mut conn).await.unwrap(), 1); + assert_eq!(count_submissions_failed(&mut conn).await.unwrap(), 0); } /// Removes the given top-level key from a JSON object, panicking if it was not present. diff --git a/opsqueue/src/producer/client.rs b/opsqueue/src/producer/client.rs index 4f1dc92d..1632050b 100644 --- a/opsqueue/src/producer/client.rs +++ b/opsqueue/src/producer/client.rs @@ -398,7 +398,6 @@ impl InternalProducerClientError { #[cfg(test)] #[cfg(feature = "server-logic")] mod tests { - use ux::u63; use crate::{ common::{ @@ -459,7 +458,7 @@ mod tests { let count = submission::db::count_submissions(&mut conn) .await .expect("Should be OK"); - assert_eq!(count, u63::new(0)); + assert_eq!(count, 0); let submission = InsertSubmission { chunk_contents: ChunkContents::Direct { @@ -477,7 +476,7 @@ mod tests { let count = submission::db::count_submissions(&mut conn) .await .expect("Should be OK"); - assert_eq!(count, u63::new(1)); + assert_eq!(count, 1); client .insert_submission(&submission, &std::collections::HashMap::default()) @@ -495,7 +494,7 @@ mod tests { let count = submission::db::count_submissions(&mut conn) .await .expect("Should be OK"); - assert_eq!(count, u63::new(4)); + assert_eq!(count, 4); } #[sqlx::test(migrator = "crate::MIGRATOR")] diff --git a/opsqueue/src/producer/server.rs b/opsqueue/src/producer/server.rs index 74b87075..1fa4d678 100644 --- a/opsqueue/src/producer/server.rs +++ b/opsqueue/src/producer/server.rs @@ -222,7 +222,7 @@ pub struct InsertSubmissionResponse { async fn submissions_count(State(state): State) -> Result, ServerError> { let mut conn = state.pool.reader_conn().await?; let count = submission::db::count_submissions(&mut conn).await?; - Ok(Json(u64::from(count))) + Ok(Json(count)) } async fn submissions_count_completed( @@ -230,5 +230,5 @@ async fn submissions_count_completed( ) -> Result, ServerError> { let mut conn = state.pool.reader_conn().await?; let count = submission::db::count_submissions_completed(&mut conn).await?; - Ok(Json(u64::from(count))) + Ok(Json(count)) } diff --git a/opsqueue/src/prometheus.rs b/opsqueue/src/prometheus.rs index 8bc7f023..d4b308a1 100644 --- a/opsqueue/src/prometheus.rs +++ b/opsqueue/src/prometheus.rs @@ -211,9 +211,7 @@ pub fn time_delta_as_f64(td: chrono::TimeDelta) -> f64 { #[allow(clippy::cast_precision_loss)] pub async fn calculate_scaling_metrics(db_pool: &DBPools) -> anyhow::Result<()> { let mut conn = db_pool.reader_conn().await?; - let chunks_backlog_count: u64 = crate::common::chunk::db::count_chunks(&mut conn) - .await? - .into(); + let chunks_backlog_count: u64 = crate::common::chunk::db::count_chunks(&mut conn).await?; gauge!(CHUNKS_BACKLOG_GAUGE).set(chunks_backlog_count as f64); let ops_backlog_count: f64 = crate::common::chunk::db::count_ops_in_backlog_estimate(&mut conn).await?; From 8780d964703e1a1fa6b3ff93ed2b9fe941e33737 Mon Sep 17 00:00:00 2001 From: Sem Mulder Date: Wed, 22 Jul 2026 12:55:48 +0200 Subject: [PATCH 02/11] Make tests timeout properly, in preparation for showing failure logs properly --- .../python/opsqueue/producer.py | 16 +++- libs/opsqueue_python/src/errors.rs | 8 +- libs/opsqueue_python/src/producer.rs | 71 +++++------------ libs/opsqueue_python/tests/test_roundtrip.py | 79 ++++++++++++++++--- 4 files changed, 106 insertions(+), 68 deletions(-) diff --git a/libs/opsqueue_python/python/opsqueue/producer.py b/libs/opsqueue_python/python/opsqueue/producer.py index 82a877e3..dc807cc6 100644 --- a/libs/opsqueue_python/python/opsqueue/producer.py +++ b/libs/opsqueue_python/python/opsqueue/producer.py @@ -96,6 +96,7 @@ def run_submission( serialization_format: SerializationFormat = DEFAULT_SERIALIZATION_FORMAT, metadata: None | bytes = None, strategic_metadata: None | dict[str, int] = None, + timeout: float | None = None, ) -> Iterator[Any]: """ Inserts a submission into the queue, and blocks until it is completed. @@ -116,6 +117,7 @@ def run_submission( metadata=metadata, strategic_metadata=strategic_metadata, chunk_size=chunk_size, + timeout=timeout, ) return _unchunk_iterator(results_iter, serialization_format) @@ -169,6 +171,7 @@ def blocking_stream_completed_submission( submission_id: SubmissionId, *, serialization_format: SerializationFormat = DEFAULT_SERIALIZATION_FORMAT, + timeout: float | None = None, ) -> Iterator[Any]: """ Blocks until the submission is completed. @@ -181,7 +184,7 @@ def blocking_stream_completed_submission( (after retrying a consumer kept failing on one of the chunks) """ return _unchunk_iterator( - self.blocking_stream_completed_submission_chunks(submission_id), + self.blocking_stream_completed_submission_chunks(submission_id, timeout), serialization_format, ) @@ -211,6 +214,7 @@ def run_submission_chunks( metadata: None | bytes = None, strategic_metadata: None | dict[str, int] = None, chunk_size: None | int = None, + timeout: float | None = None, ) -> Iterator[bytes]: """ Inserts an already-chunked submission into the queue, and blocks until it is completed. @@ -229,7 +233,7 @@ def run_submission_chunks( strategic_metadata=strategic_metadata, chunk_size=chunk_size, ) - return self.blocking_stream_completed_submission_chunks(submission_id) + return self.blocking_stream_completed_submission_chunks(submission_id, timeout) async def async_run_submission_chunks( self, @@ -278,7 +282,9 @@ def insert_submission_chunks( ) def blocking_stream_completed_submission_chunks( - self, submission_id: SubmissionId + self, + submission_id: SubmissionId, + timeout: float | None = None, ) -> Iterator[bytes]: """ Blocks until the submission is completed, and returns an iterator that lazily @@ -289,7 +295,9 @@ def blocking_stream_completed_submission_chunks( - `SubmissionFailedError` if the submission failed permanently (after retrying a consumer kept failing on one of the chunks) """ - return self.inner.blocking_stream_completed_submission_chunks(submission_id) # type: ignore[no-any-return] + return self.inner.blocking_stream_completed_submission_chunks( # type: ignore[no-any-return] + submission_id, timeout + ) async def async_stream_completed_submission_chunks( self, submission_id: SubmissionId diff --git a/libs/opsqueue_python/src/errors.rs b/libs/opsqueue_python/src/errors.rs index 45f0f7da..110dded0 100644 --- a/libs/opsqueue_python/src/errors.rs +++ b/libs/opsqueue_python/src/errors.rs @@ -7,7 +7,7 @@ use opsqueue::common::errors::{ ChunkNotFound, E, IncorrectUsage, SubmissionNotCancellable, SubmissionNotFound, TooManyMatchingSubmissions, UnexpectedOpsqueueConsumerServerResponse, }; -use pyo3::exceptions::PyBaseException; +use pyo3::exceptions::{PyBaseException, PyTimeoutError}; use pyo3::{Bound, PyErr, Python, import_exception}; use crate::common; @@ -201,6 +201,12 @@ impl From> for PyErr { } } +impl From> for PyErr { + fn from(_value: CError) -> Self { + PyTimeoutError::new_err("timeout was reached") + } +} + impl From for CError> { fn from(value: PyErr) -> Self { CError(E::L(FatalPythonException(value))) diff --git a/libs/opsqueue_python/src/producer.rs b/libs/opsqueue_python/src/producer.rs index e3167a9b..1b4071a3 100644 --- a/libs/opsqueue_python/src/producer.rs +++ b/libs/opsqueue_python/src/producer.rs @@ -18,6 +18,7 @@ use opsqueue::{ producer::client::{Client as ActualClient, InternalProducerClientError}, tracing::CarrierMap, }; +use tokio::time::error::Elapsed; use ux::u63; use crate::{ @@ -376,57 +377,6 @@ impl ProducerClient { }) } - #[pyo3(signature = (chunk_contents, metadata=None, strategic_metadata=None, chunk_size=None, otel_trace_carrier=CarrierMap::default()))] - #[allow(clippy::result_large_err, clippy::type_complexity)] - /// Submit chunks and then stream the completed output chunks. - /// - /// # Errors - /// - /// Returns an error if upload, submission creation, or streaming fails. - pub fn run_submission_chunks( - &self, - py: Python<'_>, - chunk_contents: Py, - metadata: Option, - strategic_metadata: Option, - chunk_size: Option, - otel_trace_carrier: CarrierMap, - ) -> CPyResult< - PyChunksIter, - E![ - FatalPythonException, - errors::SubmissionFailed, - ChunksStorageError, - InternalProducerClientError, - ], - > { - let submission_id = self - .insert_submission_chunks( - py, - chunk_contents, - metadata, - strategic_metadata, - chunk_size, - otel_trace_carrier, - ) - .map_err(|CError(e)| { - CError(match e { - L(e) => L(e), - R(e) => R(R(e)), - }) - })?; - let res = self - .blocking_stream_completed_submission_chunks(py, submission_id) - .map_err(|CError(e)| { - CError(match e { - L(e) => L(e), - R(L(e)) => R(L(e)), - R(R(e)) => R(R(R(e))), - }) - })?; - Ok(res) - } - /// Blocks (and short-polls) until the submission is completed. /// /// We start with a small short-polling interval @@ -442,17 +392,34 @@ impl ProducerClient { &self, py: Python<'_>, submission_id: SubmissionId, + timeout: Option, ) -> CPyResult< PyChunksIter, E![ FatalPythonException, + Elapsed, errors::SubmissionFailed, InternalProducerClientError ], > { py.detach(|| { self.block_unless_interrupted(async move { - self.stream_completed_submission_chunks(submission_id).await + let fut = self.stream_completed_submission_chunks(submission_id); + match timeout { + Some(duration) => tokio::time::timeout(Duration::from_secs_f64(duration), fut) + .await + .map_err(|err| CError(R(L(err)))) + .and_then(|err| { + err.map_err(|err| match err.0 { + L(err) => CError(L(err)), + R(err) => CError(R(R(err))), + }) + }), + None => fut.await.map_err(|err| match err.0 { + L(err) => CError(L(err)), + R(err) => CError(R(R(err))), + }), + } }) }) } diff --git a/libs/opsqueue_python/tests/test_roundtrip.py b/libs/opsqueue_python/tests/test_roundtrip.py index 23d9a149..6c86e2ce 100644 --- a/libs/opsqueue_python/tests/test_roundtrip.py +++ b/libs/opsqueue_python/tests/test_roundtrip.py @@ -29,6 +29,8 @@ import logging import pytest +SUBMISSION_COMPLETED_TIMEOUT = 10.0 + def increment(data: int) -> int: return data + 1 @@ -56,7 +58,10 @@ def run_consumer() -> None: input_iter = range(0, 100) output_iter: Iterator[int] = producer_client.run_submission( - input_iter, chunk_size=20, strategic_metadata={"id": 42} + input_iter, + chunk_size=20, + strategic_metadata={"id": 42}, + timeout=SUBMISSION_COMPLETED_TIMEOUT, ) res = sum(output_iter) @@ -128,6 +133,7 @@ def run_consumer(_consumer_id: int) -> None: input_iter, chunk_size=chunk_size, strategic_metadata={"id": 42, "second_id": 69}, + timeout=SUBMISSION_COMPLETED_TIMEOUT, ) res = sum(output_iter) @@ -146,7 +152,9 @@ def test_empty_submission(opsqueue: OpsqueueProcess) -> None: input_iter: list[int] = [] output_iter: Iterator[int] = producer_client.run_submission( - input_iter, chunk_size=20 + input_iter, + chunk_size=20, + timeout=SUBMISSION_COMPLETED_TIMEOUT, ) res = sum(output_iter) assert res == 0 @@ -182,7 +190,10 @@ def run_consumer() -> None: input_iter = range(0, 100) output_iter: Iterator[int] = producer_client.run_submission( - input_iter, chunk_size=20, serialization_format=serialization_format + input_iter, + chunk_size=20, + serialization_format=serialization_format, + timeout=SUBMISSION_COMPLETED_TIMEOUT, ) res = sum(output_iter) @@ -225,7 +236,11 @@ def broken_increment(input: int) -> float: input_iter = range(0, 100) with pytest.raises(SubmissionFailedError) as exc_info: - producer_client.run_submission(input_iter, chunk_size=20) + producer_client.run_submission( + input_iter, + chunk_size=20, + timeout=SUBMISSION_COMPLETED_TIMEOUT, + ) # We expect the intended attributes to be there: assert isinstance(exc_info.value.failure, str) @@ -265,7 +280,10 @@ def increment_list(ints: Sequence[int], _chunk: Chunk) -> Sequence[int]: input_iter = map(lambda i: cbor2.dumps([i, i, i]), range(0, 10)) output_iter: Iterator[list[int]] = map( lambda c: cbor2.loads(c), - producer_client.run_submission_chunks(input_iter), + producer_client.run_submission_chunks( + input_iter, + timeout=SUBMISSION_COMPLETED_TIMEOUT, + ), ) import itertools @@ -304,7 +322,9 @@ def run_consumer(consumer_id: int) -> None: with multiple_background_processes(run_consumer, n_consumers) as _consumers: input_iter = range(0, 1000) output_iter: Iterator[int] = producer_client.run_submission( - input_iter, chunk_size=100 + input_iter, + chunk_size=100, + timeout=SUBMISSION_COMPLETED_TIMEOUT, ) res = sum(output_iter) @@ -379,7 +399,10 @@ def run_consumer() -> None: with background_process(run_consumer): # Wait for the submission to complete. - producer_client.blocking_stream_completed_submission(submission_id) + producer_client.blocking_stream_completed_submission( + submission_id, + timeout=SUBMISSION_COMPLETED_TIMEOUT, + ) submission = producer_client.get_submission_status(submission_id) assert submission is not None assert isinstance(submission.submission, SubmissionCompleted) @@ -423,7 +446,10 @@ def assert_submission_failed_has_metadata(x: SubmissionFailed) -> None: with pytest.raises(SubmissionFailedError) as exc_info: # Wait for the submission to fail. - producer_client.blocking_stream_completed_submission(submission_id) + producer_client.blocking_stream_completed_submission( + submission_id, + timeout=SUBMISSION_COMPLETED_TIMEOUT, + ) assert_submission_failed_has_metadata(exc_info.value.submission) submission = producer_client.get_submission_status(submission_id) @@ -511,7 +537,10 @@ def run_consumer() -> None: with background_process(run_consumer): # Wait for the submission to complete. - producer_client.blocking_stream_completed_submission(submission_id) + producer_client.blocking_stream_completed_submission( + submission_id, + timeout=SUBMISSION_COMPLETED_TIMEOUT, + ) submission = producer_client.get_submission_status(submission_id) assert submission is not None assert isinstance(submission.submission, SubmissionCompleted) @@ -544,7 +573,10 @@ def consume(x: int) -> None: with background_process(run_consumer): with pytest.raises(SubmissionFailedError): - producer_client.blocking_stream_completed_submission(submission_id) + producer_client.blocking_stream_completed_submission( + submission_id, + timeout=SUBMISSION_COMPLETED_TIMEOUT, + ) # Cancelling the failed submission should fail. with pytest.raises(SubmissionNotCancellableError) as exc_info: producer_client.cancel_submission(submission_id) @@ -576,7 +608,10 @@ def consume(x: int) -> int | None: with background_process(run_consumer): with pytest.raises(SubmissionFailedError) as exc_info: - producer_client.blocking_stream_completed_submission(submission_id) + producer_client.blocking_stream_completed_submission( + submission_id, + timeout=SUBMISSION_COMPLETED_TIMEOUT, + ) assert exc_info.value.submission.chunks_done == len(chunks) - 1 @@ -667,3 +702,25 @@ def test_lookup_too_many_submission_ids_by_strategic_metadata() -> None: ) assert exc.type is TooManyMatchingSubmissionsError assert exc.value.max_submissions == max_ + + +def test_run_submission_timeout(opsqueue: OpsqueueProcess) -> None: + url = "file:///tmp/opsqueue/test_run_submission_timeout" + producer_client = ProducerClient(f"localhost:{opsqueue.port}", url) + + def run_consumer() -> None: + consumer_client = ConsumerClient(f"localhost:{opsqueue.port}", url) + + def process_op(x: int) -> int: + time.sleep(2.0) + return x + + consumer_client.run_each_op(process_op) + + with background_process(run_consumer) as _consumer: + with pytest.raises(TimeoutError): + producer_client.run_submission( + [1], + chunk_size=1, + timeout=0.1, + ) From 77c8b11c26f6f029e3f0515f91d43616c5186fa2 Mon Sep 17 00:00:00 2001 From: Sem Mulder Date: Thu, 23 Jul 2026 16:16:43 +0200 Subject: [PATCH 03/11] Make complete_chunk and fail_chunk not error when processing previously completed, failed, or cancelled chunks Because of the idempotency assumption for processing chunks, nothing should break if we just ignore the error. Besides, we were already ignoring the error accidentally. --- .../python/opsqueue/exceptions.py | 9 -- libs/opsqueue_python/src/errors.rs | 23 +--- opsqueue/src/common/chunk.rs | 127 ++++++++++-------- opsqueue/src/common/errors.rs | 6 +- 4 files changed, 76 insertions(+), 89 deletions(-) diff --git a/libs/opsqueue_python/python/opsqueue/exceptions.py b/libs/opsqueue_python/python/opsqueue/exceptions.py index c946f339..542aa7a6 100644 --- a/libs/opsqueue_python/python/opsqueue/exceptions.py +++ b/libs/opsqueue_python/python/opsqueue/exceptions.py @@ -92,15 +92,6 @@ class TryFromIntError(IncorrectUsageError): pass -class ChunkNotFoundError(IncorrectUsageError): - """ - Raised when a method is used to look up information about a chunk - but the chunk doesn't exist within the Opsqueue. - """ - - pass - - class SubmissionNotFoundError(IncorrectUsageError): """ Raised when a method is used to look up information about a submission diff --git a/libs/opsqueue_python/src/errors.rs b/libs/opsqueue_python/src/errors.rs index 110dded0..b0de8e7e 100644 --- a/libs/opsqueue_python/src/errors.rs +++ b/libs/opsqueue_python/src/errors.rs @@ -2,16 +2,14 @@ /// so we have nice IDE support for docs-on-hover and for 'go to definition'. use std::error::Error; -use opsqueue::common::chunk::ChunkId; use opsqueue::common::errors::{ - ChunkNotFound, E, IncorrectUsage, SubmissionNotCancellable, SubmissionNotFound, - TooManyMatchingSubmissions, UnexpectedOpsqueueConsumerServerResponse, + E, IncorrectUsage, SubmissionNotCancellable, SubmissionNotFound, TooManyMatchingSubmissions, + UnexpectedOpsqueueConsumerServerResponse, }; use pyo3::exceptions::{PyBaseException, PyTimeoutError}; use pyo3::{Bound, PyErr, Python, import_exception}; use crate::common; -use crate::common::{ChunkIndex, SubmissionId}; // Expected errors: import_exception!(opsqueue.exceptions, SubmissionFailedError); @@ -19,7 +17,6 @@ import_exception!(opsqueue.exceptions, SubmissionFailedError); // Incorrect usage errors: import_exception!(opsqueue.exceptions, IncorrectUsageError); import_exception!(opsqueue.exceptions, TryFromIntError); -import_exception!(opsqueue.exceptions, ChunkNotFoundError); import_exception!(opsqueue.exceptions, SubmissionNotFoundError); import_exception!(opsqueue.exceptions, SubmissionNotCancellableError); import_exception!(opsqueue.exceptions, TooManyMatchingSubmissionsError); @@ -173,22 +170,6 @@ impl From> for PyErr { } } -impl From> for PyErr { - fn from(value: CError) -> Self { - let ChunkId { - submission_id, - chunk_index, - } = value.0.0; - ChunkNotFoundError::new_err(( - value.0.to_string(), - ( - SubmissionId::from(submission_id), - ChunkIndex::from(chunk_index), - ), - )) - } -} - impl From> for PyErr { fn from(value: CError) -> Self { NewObjectStoreClientError::new_err(value.0.to_string()) diff --git a/opsqueue/src/common/chunk.rs b/opsqueue/src/common/chunk.rs index a56515d5..ba60b74a 100644 --- a/opsqueue/src/common/chunk.rs +++ b/opsqueue/src/common/chunk.rs @@ -225,13 +225,13 @@ impl Chunk { #[cfg(feature = "server-logic")] pub mod db { use super::{ - Chunk, ChunkCompleted, ChunkFailed, ChunkId, ChunkIndex, ChunkSize, DateTime, SubmissionId, + Chunk, ChunkCompleted, ChunkFailed, ChunkId, ChunkIndex, DateTime, SubmissionId, Utc, }; - use crate::common::errors::{ChunkNotFound, DatabaseError, E, SubmissionNotFound}; + use crate::common::errors::{DatabaseError, E, SubmissionNotFound}; use crate::db::{Connection, True, WriterConnection}; use axum_prometheus::metrics::{counter, gauge}; use sqlx::{QueryBuilder, Sqlite}; - use sqlx::{query, query_as}; + use sqlx::{query, query_as, query_scalar}; impl<'q> sqlx::Encode<'q, Sqlite> for super::ChunkIndex { fn encode_by_ref( @@ -300,25 +300,18 @@ pub mod db { chunk_id: ChunkId, output_content: Option>, mut conn: impl WriterConnection, - ) -> Result<(), E>> { - let _chunk_size: Result>> = - conn.transaction(move |mut tx| { - Box::pin(async move { - let completed_work = - complete_chunk_raw(chunk_id, output_content, &mut tx).await?; - crate::common::submission::db::maybe_complete_submission( - chunk_id.submission_id, - &mut tx, - ) - .await - .map_err(|e| match e { - E::L(e) => E::L(e), - E::R(e) => E::R(E::L(e)), - })?; - Ok(completed_work.unwrap_or_default()) - }) + ) -> Result<(), E> { + conn.transaction(move |mut tx| { + Box::pin(async move { + complete_chunk_raw(chunk_id, output_content, &mut tx).await?; + crate::common::submission::db::maybe_complete_submission( + chunk_id.submission_id, + &mut tx, + ) + .await }) - .await; + }) + .await?; counter!(crate::prometheus::CHUNKS_COMPLETED_COUNTER).increment(1); Ok(()) @@ -334,9 +327,9 @@ pub mod db { chunk_id: ChunkId, output_content: Option>, mut tx: impl WriterConnection, - ) -> sqlx::Result> { + ) -> sqlx::Result<()> { let now = chrono::prelude::Utc::now(); - query!( + let chunk_moved = query!( " INSERT INTO chunks_completed (submission_id, chunk_index, output_content, completed_at) @@ -353,26 +346,42 @@ pub mod db { chunk_id.submission_id, chunk_id.chunk_index, ) - .fetch_one(tx.get_inner()) - .await?; - // Defense in depth: Above query should never be called twice on the same chunk. - // If it _does_ happen, it means that either a consumer is attempting a chunk they didn't reserve, - // or we gave out the same reservation twice. + .fetch_optional(tx.get_inner()) + .await? + .is_some(); + // Defense in depth: Above query could be called twice on the same chunk. For instance, + // when the server was restarted and the reservations are forgotten, and the same chunk + // was reserved again. + // + // In addition, cancelling a submission while a chunk is reserved also results in the chunk + // not being in the `chunks` table. Which is fine, because cancelled submissions count as + // failed. + // + // By only updating `chunks_done` when we actually moved a chunk, we ensure that we never + // mess up the submission's `chunks_done` counter. // - // By returning early if the chunk was not found, - // we ensure that even in these situations - // we never mess up the submission's `chunks_done` counter. + // This does mean we potentially run the same chunk twice, but that is fine because we + // assume chunks to be processed idempotently. // // (Not doing that resulted in a hard-to-track-down bug in the past. // https://github.com/channable/opsqueue/issues/76 // ) - sqlx::query_scalar!( - "UPDATE submissions SET chunks_done = chunks_done + 1 WHERE submissions.id = $1 RETURNING submissions.chunk_size;", - chunk_id.submission_id, - ) - .fetch_one(tx.get_inner()) - .await - .map(|opt| opt.map(ChunkSize)) + if chunk_moved { + sqlx::query_scalar!( + "UPDATE submissions SET chunks_done = chunks_done + 1 WHERE submissions.id = $1 RETURNING submissions.chunk_size;", + chunk_id.submission_id, + ) + .fetch_one(tx.get_inner()) + .await?; + } else { + tracing::warn!( + "Could not complete chunk {:?} because it was either: \ + completed, failed, or cancelled before. Ignoring.", + chunk_id + ); + } + + Ok(()) } /// Increment retries for a chunk, or move it to failed state. @@ -394,7 +403,7 @@ pub mod db { submission_id, chunk_index, } = chunk_id; - let fields = query!( + let retries = query_scalar!( " UPDATE chunks SET retries = retries + 1 WHERE submission_id = $1 AND chunk_index = $2 @@ -403,23 +412,33 @@ pub mod db { submission_id, chunk_index ) - .fetch_one(tx.get_inner()) + .fetch_optional(tx.get_inner()) .await?; - tracing::trace!("Retries: {}", fields.retries); - if fields.retries >= max_retries.into() { - crate::common::submission::db::fail_submission_notx( - submission_id, - chunk_index, - failure, - &mut tx, - ) - .await?; - - Ok::<_, sqlx::Error>(true) + if let Some(retries) = retries { + tracing::trace!("Retries: {}", retries); + if retries >= max_retries.into() { + crate::common::submission::db::fail_submission_notx( + submission_id, + chunk_index, + failure, + &mut tx, + ) + .await?; + + Ok::<_, sqlx::Error>(true) + } else { + counter!(crate::prometheus::CHUNKS_RETRIED_COUNTER).increment(1); + // When retrying, the chunk re-enters ('stays') in the backlog, + // so we *don't* decrement the backlog gauge here. + Ok::<_, sqlx::Error>(false) + } } else { - counter!(crate::prometheus::CHUNKS_RETRIED_COUNTER).increment(1); - // When retrying, the chunk re-enters ('stays') in the backlog, - // so we *don't* decrement the backlog gauge here. + tracing::warn!( + "Could not fail chunk {:?} because it was either: \ + completed, failed, or cancelled before. Ignoring.", + chunk_id + ); + Ok::<_, sqlx::Error>(false) } }) diff --git a/opsqueue/src/common/errors.rs b/opsqueue/src/common/errors.rs index 6527810b..6504ecd5 100644 --- a/opsqueue/src/common/errors.rs +++ b/opsqueue/src/common/errors.rs @@ -12,7 +12,7 @@ use thiserror::Error; use crate::consumer::common::SyncServerToClientResponse; use super::{ - chunk::{ChunkFailed, ChunkId}, + chunk::ChunkFailed, submission::{SubmissionCancelled, SubmissionCompleted, SubmissionFailed, SubmissionId}, }; @@ -28,10 +28,6 @@ impl From for E { } } -#[derive(Error, Debug)] -#[error("Chunk not found for ID {0:?}")] -pub struct ChunkNotFound(pub ChunkId); - #[derive(Error, Debug, Deserialize, Serialize)] #[error("Submission not found for ID {0:?}")] pub struct SubmissionNotFound(pub SubmissionId); From 4d59aa2d49b32da468cfd734fb7749da62fe529c Mon Sep 17 00:00:00 2001 From: Sem Mulder Date: Thu, 23 Jul 2026 16:32:15 +0200 Subject: [PATCH 04/11] Allow submissions to be created in a paused state Introduce `submissions_paused` and `chunks_paused` tables (alongside the existing `submissions_{completed,failed,cancelled}` and `chunks_{completed,failed}` tables). A submission can now be created in a Paused state. It's then stored in `submissions_paused` and its chunks are stored in `chunks_paused`. Because paused chunks are not in the `chunks` table, the consumer dispatcher naturally skips them without any changes to the dispatch query. Unpausing moves the submission and the chunks to `submissions` and `chunks` and notifies waiting consumers. Paused submissions are cancellable; `cancel_submission` now handles the case where the submission is found in `submissions_paused`. We don't allow pausing submissions after creation. That proved to have too many edge cases we would need to resolve. --- .../python/opsqueue/producer.py | 21 +- libs/opsqueue_python/src/common.rs | 44 +- libs/opsqueue_python/src/lib.rs | 1 + libs/opsqueue_python/src/producer.rs | 40 +- libs/opsqueue_python/tests/test_roundtrip.py | 66 +++ .../20260805143000_pausing.down.sql | 2 + .../migrations/20260805143000_pausing.up.sql | 22 + opsqueue/opsqueue_example_database_schema.db | Bin 102400 -> 106496 bytes opsqueue/src/common/chunk.rs | 115 ++++- opsqueue/src/common/submission.rs | 447 +++++++++++++++++- opsqueue/src/consumer/client.rs | 1 + opsqueue/src/consumer/strategy.rs | 1 + opsqueue/src/producer/client.rs | 125 ++++- opsqueue/src/producer/common.rs | 4 + opsqueue/src/producer/server.rs | 34 +- opsqueue/src/prometheus.rs | 12 + 16 files changed, 896 insertions(+), 39 deletions(-) create mode 100644 opsqueue/migrations/20260805143000_pausing.down.sql create mode 100644 opsqueue/migrations/20260805143000_pausing.up.sql diff --git a/libs/opsqueue_python/python/opsqueue/producer.py b/libs/opsqueue_python/python/opsqueue/producer.py index dc807cc6..9bb72d6c 100644 --- a/libs/opsqueue_python/python/opsqueue/producer.py +++ b/libs/opsqueue_python/python/opsqueue/producer.py @@ -27,6 +27,7 @@ SubmissionFailed, ChunkFailed, SubmissionNotCancellable, + SubmissionPaused, ) __all__ = [ @@ -39,6 +40,7 @@ "SubmissionNotCancellable", "SubmissionNotCancellableError", "SubmissionNotFoundError", + "SubmissionPaused", "TooManyMatchingSubmissionsError", "ChunkFailed", ] @@ -148,6 +150,7 @@ def insert_submission( serialization_format: SerializationFormat = DEFAULT_SERIALIZATION_FORMAT, metadata: None | bytes = None, strategic_metadata: None | dict[str, int] = None, + paused: bool = False, ) -> SubmissionId: """ Inserts a submission into the queue, @@ -164,6 +167,7 @@ def insert_submission( metadata=metadata, strategic_metadata=strategic_metadata, chunk_size=chunk_size, + paused=paused, ) def blocking_stream_completed_submission( @@ -263,6 +267,7 @@ def insert_submission_chunks( metadata: None | bytes = None, strategic_metadata: None | dict[str, int] = None, chunk_size: None | int = None, + paused: bool = False, ) -> SubmissionId: """ Inserts an already-chunked submission into the queue, @@ -279,6 +284,7 @@ def insert_submission_chunks( strategic_metadata=strategic_metadata, chunk_size=chunk_size, otel_trace_carrier=otel_trace_carrier, + paused=paused, ) def blocking_stream_completed_submission_chunks( @@ -334,7 +340,7 @@ def count_submissions(self) -> int: def cancel_submission(self, submission_id: SubmissionId) -> None: """ - Cancel a specific submission that is in progress. + Cancel a specific submission that is in progress or paused. Returns None if the submission was successfully cancelled. @@ -345,6 +351,19 @@ def cancel_submission(self, submission_id: SubmissionId) -> None: """ self.inner.cancel_submission(submission_id) + def unpause_submission(self, submission_id: SubmissionId) -> None: + """ + Unpause a specific submission that is currently paused, + making it available to consumers. + + Returns None if the submission was successfully unpaused. + + Raises: + - `SubmissionNotFoundError` if the submission is not currently paused. + - `InternalProducerClientError` if there is a low-level internal error. + """ + self.inner.unpause_submission(submission_id) + def get_submission_status( self, submission_id: SubmissionId ) -> SubmissionStatus | None: diff --git a/libs/opsqueue_python/src/common.rs b/libs/opsqueue_python/src/common.rs index b0fbc0b3..abe6b535 100644 --- a/libs/opsqueue_python/src/common.rs +++ b/libs/opsqueue_python/src/common.rs @@ -364,12 +364,15 @@ pub enum SubmissionStatus { Cancelled { submission: SubmissionCancelled, }, + Paused { + submission: SubmissionPaused, + }, } impl From for SubmissionStatus { fn from(value: opsqueue::common::submission::SubmissionStatus) -> Self { use opsqueue::common::submission::SubmissionStatus::{ - Cancelled, Completed, Failed, InProgress, + Cancelled, Completed, Failed, InProgress, Paused, }; match value { InProgress(s) => SubmissionStatus::InProgress { @@ -386,6 +389,9 @@ impl From for SubmissionStatus { Cancelled(s) => SubmissionStatus::Cancelled { submission: s.into(), }, + Paused(s) => SubmissionStatus::Paused { + submission: s.into(), + }, } } } @@ -510,6 +516,42 @@ pub struct SubmissionCancelled { pub cancelled_at: DateTime, } +#[pyclass(from_py_object, frozen, get_all, module = "opsqueue")] +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SubmissionPaused { + pub id: SubmissionId, + pub chunks_total: u64, + pub chunks_done: u64, + pub metadata: Option, + pub strategic_metadata: StrategicMetadataMap, +} + +impl From for SubmissionPaused { + fn from(value: opsqueue::common::submission::SubmissionPaused) -> Self { + Self { + id: value.id.into(), + chunks_total: value.chunks_total.into(), + chunks_done: value.chunks_done.into(), + metadata: value.metadata, + strategic_metadata: value.strategic_metadata, + } + } +} + +#[pymethods] +impl SubmissionPaused { + fn __repr__(&self) -> String { + format!( + "SubmissionPaused(id={0}, chunks_total={1}, chunks_done={2}, metadata={3:?}, strategic_metadata={4:?})", + self.id.__repr__(), + self.chunks_total, + self.chunks_done, + self.metadata, + self.strategic_metadata + ) + } +} + /// Submission could not be cancelled because it was already completed, failed /// or cancelled. #[pyclass(from_py_object, frozen, module = "opsqueue")] diff --git a/libs/opsqueue_python/src/lib.rs b/libs/opsqueue_python/src/lib.rs index b5f804f5..27f6835c 100644 --- a/libs/opsqueue_python/src/lib.rs +++ b/libs/opsqueue_python/src/lib.rs @@ -24,6 +24,7 @@ fn opsqueue_internal(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_class::()?; m.add_class::()?; m.add_class::()?; + m.add_class::()?; m.add_class::()?; m.add_class::()?; m.add_class::()?; diff --git a/libs/opsqueue_python/src/producer.rs b/libs/opsqueue_python/src/producer.rs index 1b4071a3..1c40ad25 100644 --- a/libs/opsqueue_python/src/producer.rs +++ b/libs/opsqueue_python/src/producer.rs @@ -159,6 +159,36 @@ impl ProducerClient { }) } + /// Unpause a paused submission, making it available to consumers again. + /// + /// Will return an error if the submission is not currently paused. + /// + /// # Errors + /// + /// Returns an error if the submission is not found or if an internal client error occurs. + #[allow(clippy::result_large_err, clippy::type_complexity)] + pub fn unpause_submission( + &self, + py: Python<'_>, + id: SubmissionId, + ) -> CPyResult< + (), + E![ + FatalPythonException, + SubmissionNotFound, + InternalProducerClientError + ], + > { + py.detach(|| { + self.block_unless_interrupted(async { + self.client + .unpause_submission(id.into()) + .await + .map_err(|e| CError(R(e))) + }) + }) + } + /// Retrieve the status (in progress, completed or failed) of a specific submission. /// /// The returned `SubmissionStatus` object also includes the number of chunks finished so far, @@ -247,7 +277,7 @@ impl ProducerClient { /// # Errors /// /// Returns an error if submission insertion fails. - #[pyo3(signature = (chunk_contents, metadata=None, chunk_size=None, otel_trace_carrier=CarrierMap::default()))] + #[pyo3(signature = (chunk_contents, metadata=None, chunk_size=None, otel_trace_carrier=CarrierMap::default(), paused=false))] pub fn insert_submission_direct( &self, py: Python<'_>, @@ -255,6 +285,7 @@ impl ProducerClient { metadata: Option, chunk_size: Option, otel_trace_carrier: CarrierMap, + paused: bool, ) -> CPyResult> { let strategic_metadata = std::collections::HashMap::default(); @@ -266,6 +297,7 @@ impl ProducerClient { }, metadata, strategic_metadata, + paused, }; self.block_unless_interrupted(async move { self.client @@ -277,8 +309,8 @@ impl ProducerClient { }) } - #[pyo3(signature = (chunk_contents, metadata=None, strategic_metadata=None, chunk_size=None, otel_trace_carrier=CarrierMap::default()))] - #[allow(clippy::type_complexity)] + #[pyo3(signature = (chunk_contents, metadata=None, strategic_metadata=None, chunk_size=None, otel_trace_carrier=CarrierMap::default(), paused=false))] + #[allow(clippy::type_complexity, clippy::too_many_arguments)] /// Insert submission chunks via object storage and enqueue the submission. /// /// # Errors @@ -292,6 +324,7 @@ impl ProducerClient { strategic_metadata: Option, chunk_size: Option, otel_trace_carrier: CarrierMap, + paused: bool, ) -> CPyResult< SubmissionId, E![ @@ -331,6 +364,7 @@ impl ProducerClient { }, metadata, strategic_metadata: strategic_metadata.unwrap_or_default(), + paused, }; self.client .insert_submission(&submission, &otel_trace_carrier) diff --git a/libs/opsqueue_python/tests/test_roundtrip.py b/libs/opsqueue_python/tests/test_roundtrip.py index 6c86e2ce..6279ccc0 100644 --- a/libs/opsqueue_python/tests/test_roundtrip.py +++ b/libs/opsqueue_python/tests/test_roundtrip.py @@ -27,6 +27,7 @@ strategy_from_description, ) import logging +import time import pytest SUBMISSION_COMPLETED_TIMEOUT = 10.0 @@ -724,3 +725,68 @@ def process_op(x: int) -> int: chunk_size=1, timeout=0.1, ) + + +def test_unpause_and_complete(opsqueue: OpsqueueProcess) -> None: + """Unpausing a paused submission makes it available to consumers again, + and it can be completed normally afterwards.""" + url = "file:///tmp/opsqueue/test_unpause_and_complete" + producer_client = ProducerClient(f"localhost:{opsqueue.port}", url) + submission_id = producer_client.insert_submission( + (1, 2, 3), chunk_size=1, paused=True + ) + + assert isinstance( + producer_client.get_submission_status(submission_id), SubmissionStatus.Paused + ) + + producer_client.unpause_submission(submission_id) + assert isinstance( + producer_client.get_submission_status(submission_id), + SubmissionStatus.InProgress, + ) + + def run_consumer() -> None: + consumer_client = ConsumerClient(f"localhost:{opsqueue.port}", url) + consumer_client.run_each_op(lambda x: x) + + with background_process(run_consumer): + producer_client.blocking_stream_completed_submission(submission_id) + assert isinstance( + producer_client.get_submission_status(submission_id), + SubmissionStatus.Completed, + ) + + +def test_unpause_not_found(opsqueue: OpsqueueProcess) -> None: + """Unpausing a submission that is not paused (e.g. in-progress) raises + SubmissionNotFoundError.""" + url = "file:///tmp/opsqueue/test_unpause_not_found" + producer_client = ProducerClient(f"localhost:{opsqueue.port}", url) + submission_id = producer_client.insert_submission( + (1, 2, 3), chunk_size=1, paused=False + ) + assert isinstance( + producer_client.get_submission_status(submission_id), + SubmissionStatus.InProgress, + ) + with pytest.raises(SubmissionNotFoundError): + producer_client.unpause_submission(submission_id) + + +def test_cancel_paused(opsqueue: OpsqueueProcess) -> None: + """A paused submission can be cancelled; its status becomes Cancelled.""" + url = "file:///tmp/opsqueue/test_cancel_paused" + producer_client = ProducerClient(f"localhost:{opsqueue.port}", url) + submission_id = producer_client.insert_submission( + (1, 2, 3), chunk_size=1, paused=True + ) + + assert isinstance( + producer_client.get_submission_status(submission_id), SubmissionStatus.Paused + ) + + producer_client.cancel_submission(submission_id) + assert isinstance( + producer_client.get_submission_status(submission_id), SubmissionStatus.Cancelled + ) diff --git a/opsqueue/migrations/20260805143000_pausing.down.sql b/opsqueue/migrations/20260805143000_pausing.down.sql new file mode 100644 index 00000000..eed0c27f --- /dev/null +++ b/opsqueue/migrations/20260805143000_pausing.down.sql @@ -0,0 +1,2 @@ +DROP TABLE chunks_paused; +DROP TABLE submissions_paused; diff --git a/opsqueue/migrations/20260805143000_pausing.up.sql b/opsqueue/migrations/20260805143000_pausing.up.sql new file mode 100644 index 00000000..2a3a60f4 --- /dev/null +++ b/opsqueue/migrations/20260805143000_pausing.up.sql @@ -0,0 +1,22 @@ +CREATE TABLE submissions_paused +( + id BIGINT PRIMARY KEY NOT NULL, + prefix TEXT, + chunks_total INTEGER NOT NULL DEFAULT 0, + chunks_done INTEGER NOT NULL DEFAULT 0, + metadata BLOB, + otel_trace_carrier TEXT NOT NULL DEFAULT '{}', + chunk_size INTEGER +); + +CREATE INDEX submissions_paused_prefix ON submissions_paused (prefix, id); + +CREATE TABLE chunks_paused +( + submission_id INTEGER NOT NULL, + chunk_index INTEGER NOT NULL, + input_content BLOB NULL, + retries INTEGER NOT NULL DEFAULT 0, + + PRIMARY KEY (submission_id, chunk_index) +) WITHOUT ROWID, STRICT; diff --git a/opsqueue/opsqueue_example_database_schema.db b/opsqueue/opsqueue_example_database_schema.db index 83e09941a8e01e276a01805f785de1774aa7e3b4..3d3cdb8976d9c200f9abdf79f3fce5f15aded4b1 100644 GIT binary patch delta 598 zcmZozz}9epZGyC*Gy?;J6cEFJ?nE79M(K?SOZeq@x#lwPFW`6OyUUly`-``qSBqyC zPXPBd?i{XPTyr-o3RH8&M)R_>N*iDUq$2l5dDpi|k_+mC?;Lm1ob~gw z{GTu@MxOUNn}11hvC8rC7cuaE!|;{v-QM>E@A)@2hXe%mgk z$k@op!f}F`d;5GPMs{YGCJ|Y7admaZ=Jb-pq@2{`jMBX9;&_msQd3YkQqDoHjv=lJ zA&yQyt_mnp)6c6jYEO?8VYHggC(5WXy;O$Ljk%FgX8Kz#MhO;>S<@$IGm5f9OckHL zUYk)8EIj!?tE4+bR6(PlC^ap!LPsGpMYG9Xie22$kg-J*?5N_>q}`R!n9DFl zL?$;dNwPuAntXtF(PRf!w#g1GnhNUGwdy)t3P1pMb9`}TRjPugUx=%_YY>;Fv-IS> sJQ~Unees5v=86In2?r>lY`>tvXu!V6Kw*)Cz$6Fm?c52BU-f}(0J{CZ4gdfE delta 446 zcmZoTz}B#UZGyC*Bm)Bj2#W(TGZ1S{)G=n1+?cS0UxtV47X$wSepkM`e0jXTc>8&^ zcy{pwa9`uj;rg}NP~jHW=3i1=tTH_OwG8~<_+Rrs;J?a$ihn=<7XFp|^Y|z8cktJ4 z7Bq|R*crJtGaCGtU*sS#F@T+sX*)*( H;}3lRZ=!@V diff --git a/opsqueue/src/common/chunk.rs b/opsqueue/src/common/chunk.rs index ba60b74a..2526f5aa 100644 --- a/opsqueue/src/common/chunk.rs +++ b/opsqueue/src/common/chunk.rs @@ -601,6 +601,93 @@ pub mod db { Ok(()) } + /// # Errors + /// + /// Returns an error if a SQL query fails. + #[tracing::instrument(skip(chunks, conn))] + pub async fn insert_many_paused_chunks( + chunks: &[Chunk], + mut conn: impl WriterConnection, + ) -> sqlx::Result<()> { + const ROWS_PER_QUERY: usize = 1000; + + let mut iter = chunks.iter().peekable(); + while iter.peek().is_some() { + let query_chunks = iter.by_ref().take(ROWS_PER_QUERY); + + let mut query_builder: QueryBuilder = QueryBuilder::new( + "INSERT INTO chunks_paused (submission_id, chunk_index, input_content) ", + ); + query_builder.push_values(query_chunks, |mut b, chunk| { + b.push_bind(chunk.submission_id) + .push_bind(chunk.chunk_index) + .push_bind(chunk.input_content.clone()); + }); + let query = query_builder.build(); + + query.execute(conn.get_inner()).await?; + } + + Ok(()) + } + + /// Move all chunks of a paused submission from `chunks_paused` back to `chunks`. + /// + /// # Errors + /// + /// Returns an error if the SQL query fails. + #[tracing::instrument(skip(conn))] + pub async fn restore_paused_chunks( + submission_id: SubmissionId, + mut conn: impl WriterConnection, + ) -> sqlx::Result<()> { + sqlx::query!( + " + INSERT INTO chunks (submission_id, chunk_index, input_content, retries) + SELECT submission_id, chunk_index, input_content, retries FROM chunks_paused WHERE submission_id = $1; + + DELETE FROM chunks_paused WHERE submission_id = $2; + ", + submission_id, + submission_id, + ) + .execute(conn.get_inner()) + .await?; + Ok(()) + } + + /// Skip (cancel) all chunks of a paused submission by moving them from + /// `chunks_paused` to `chunks_failed` with `skipped = true`. + /// + /// # Errors + /// + /// Returns an error if the SQL query fails. + #[tracing::instrument(skip(conn))] + pub async fn skip_remaining_paused_chunks( + submission_id: SubmissionId, + mut conn: impl WriterConnection, + ) -> sqlx::Result<()> { + let now = chrono::prelude::Utc::now(); + + let query_res = sqlx::query!( + " + INSERT INTO chunks_failed + (submission_id, chunk_index, input_content, failure, skipped, failed_at) + SELECT submission_id, chunk_index, input_content, '', 1, julianday($1) FROM chunks_paused WHERE submission_id = $2; + + DELETE FROM chunks_paused WHERE submission_id = $3; + ", + now, + submission_id, + submission_id, + ) + .execute(conn.get_inner()) + .await?; + + counter!(crate::prometheus::CHUNKS_SKIPPED_COUNTER).increment(query_res.rows_affected()); + Ok(()) + } + /// Mark all remaining chunks for a submission as skipped/failed. /// /// # Errors @@ -684,6 +771,23 @@ pub mod db { Ok(u64::try_from(count).expect("COUNT(*) is always non-negative")) } + /// Count paused chunks. + /// + /// # Errors + /// + /// Returns an error if the count query fails. + /// + /// # Panics + /// + /// Panics if `COUNT(*)` returns a negative value, which `SQLite` never does. + #[tracing::instrument(skip(db))] + pub async fn count_chunks_paused(mut db: impl Connection) -> sqlx::Result { + let count = sqlx::query_scalar!("SELECT COUNT(1) as count FROM chunks_paused;") + .fetch_one(db.get_inner()) + .await?; + Ok(u64::try_from(count).expect("COUNT(*) is always non-negative")) + } + /// Looks up the number of operations in the backlog. /// /// An estimation that returns a slightly too high number, @@ -780,10 +884,7 @@ pub mod test { .expect("complete chunk failed"); assert_eq!(count_chunks(&mut conn).await.unwrap(), 0); - assert_eq!( - count_chunks_completed(&mut conn).await.unwrap(), - 1 - ); + assert_eq!(count_chunks_completed(&mut conn).await.unwrap(), 1); assert_eq!(count_chunks_failed(&mut conn).await.unwrap(), 0); } @@ -797,6 +898,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -852,10 +954,7 @@ pub mod test { .expect("Succeed chunk failed"); assert_eq!(count_chunks(&mut conn).await.unwrap(), 0); - assert_eq!( - count_chunks_completed(&mut conn).await.unwrap(), - 0 - ); + assert_eq!(count_chunks_completed(&mut conn).await.unwrap(), 0); assert_eq!(count_chunks_failed(&mut conn).await.unwrap(), 1); } } diff --git a/opsqueue/src/common/submission.rs b/opsqueue/src/common/submission.rs index a742b381..82fa39a9 100644 --- a/opsqueue/src/common/submission.rs +++ b/opsqueue/src/common/submission.rs @@ -212,12 +212,32 @@ pub struct SubmissionCancelled { pub cancelled_at: DateTime, } +/// A submission that has been paused. +/// +/// A submission can only be submitted in a paused state. We don't support pausing submissions +/// after submission. +/// +/// A paused submission can be unpaused or canceled. +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +pub struct SubmissionPaused { + pub id: SubmissionId, + pub prefix: Option, + pub chunks_total: ChunkCount, + pub chunks_done: ChunkCount, + pub chunk_size: ChunkSize, + pub metadata: Option, + #[serde(default)] + pub strategic_metadata: StrategicMetadataMap, + pub otel_trace_carrier: String, +} + #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] pub enum SubmissionStatus { InProgress(Submission), Completed(SubmissionCompleted), Failed(SubmissionFailed, ChunkFailed), Cancelled(SubmissionCancelled), + Paused(SubmissionPaused), } impl Default for Submission { @@ -284,6 +304,7 @@ pub mod db { DatabaseError, E, SubmissionNotCancellable, SubmissionNotFound, TooManyMatchingSubmissions, }, + submission::SubmissionPaused, }, db::{Connection, True, WriterConnection, WriterPool}, }; @@ -429,9 +450,123 @@ pub mod db { res } + #[tracing::instrument(skip(chunks, conn))] + pub(crate) async fn insert_paused_submission( + submission: Submission, + chunks: Vec, + mut conn: impl WriterConnection, + ) -> Result<(), DatabaseError> { + use axum_prometheus::metrics::counter; + use futures::FutureExt as _; + + let chunks_total = submission.chunks_total.into(); + tracing::debug!("Inserting paused submission {}", submission.id); + + let res = conn + .transaction(move |mut tx| { + async move { + insert_paused_submission_raw(&submission, &mut tx).await?; + insert_submission_metadata_raw( + &submission, + &submission.strategic_metadata, + &mut tx, + ) + .await?; + super::chunk::db::insert_many_paused_chunks(&chunks, &mut tx).await?; + Ok(()) + } + .boxed() + }) + .await; + + counter!(crate::prometheus::SUBMISSIONS_PAUSED_COUNTER).increment(1); + counter!(crate::prometheus::SUBMISSIONS_TOTAL_COUNTER).increment(1); + counter!(crate::prometheus::CHUNKS_TOTAL_COUNTER).increment(chunks_total); + res + } + + #[tracing::instrument(skip(conn))] + async fn insert_paused_submission_raw( + submission: &Submission, + mut conn: impl WriterConnection, + ) -> Result<(), DatabaseError> { + sqlx::query!( + " + INSERT INTO submissions_paused (id, prefix, chunks_total, chunks_done, metadata, otel_trace_carrier, chunk_size) + VALUES ($1, $2, $3, $4, $5, $6, $7) + ", + submission.id, + submission.prefix, + submission.chunks_total, + submission.chunks_done, + submission.metadata, + submission.otel_trace_carrier, + submission.chunk_size.0, + ) + .execute(conn.get_inner()) + .await?; + + Ok(()) + } + + /// Unpause a paused submission. Atomically moves it back from `submissions_paused` + /// to `submissions` and its chunks from `chunks_paused` to `chunks`. + /// + /// # Errors + /// + /// Returns [`DatabaseError`] if the transaction or any SQL query fails. + /// + /// Returns [`SubmissionNotFound`] if the submission is not currently paused. + #[tracing::instrument(skip(conn))] + pub async fn unpause_submission( + id: SubmissionId, + mut conn: impl WriterConnection, + ) -> Result<(), E> { + conn.transaction(move |mut tx| { + Box::pin(async move { + unpause_submission_raw(id, &mut tx).await?; + super::chunk::db::restore_paused_chunks(id, &mut tx).await?; + Ok(()) + }) + }) + .await + } + + #[tracing::instrument(skip(conn))] + pub(super) async fn unpause_submission_raw( + id: SubmissionId, + mut conn: impl WriterConnection, + ) -> Result<(), E> { + let row = query!( + " + INSERT INTO submissions + (id, chunks_total, chunks_done, prefix, metadata, otel_trace_carrier, chunk_size) + SELECT id, chunks_total, chunks_done, prefix, metadata, otel_trace_carrier, chunk_size + FROM submissions_paused WHERE id = $1; + + DELETE FROM submissions_paused WHERE id = $2 RETURNING *; + ", + id, + id, + ) + .fetch_optional(conn.get_inner()) + .await?; + if row.is_none() { + Err(E::R(SubmissionNotFound(id))) + } else { + counter!(crate::prometheus::SUBMISSIONS_UNPAUSED_COUNTER).increment(1); + Ok(()) + } + } + /// Creates a new submission with the given chunks and inserts it into the database. /// - /// If the number of chunks is 0, the submission is marked as completed immediately afterwards. + /// If `paused` is false and the number of chunks is 0, the submission is marked + /// as completed immediately afterwards. + /// + /// If `paused` is true, the submission is inserted directly into `submissions_paused` + /// (and its chunks into `chunks_paused`), so it won't be picked up by consumers + /// until explicitly unpaused. Zero-chunk paused submissions stay paused. /// /// # Panics /// @@ -447,6 +582,7 @@ pub mod db { metadata: Option, strategic_metadata: StrategicMetadataMap, chunk_size: ChunkSize, + paused: bool, mut conn: impl WriterConnection, ) -> Result { let submission_id = SubmissionId::new(); @@ -462,7 +598,7 @@ pub mod db { strategic_metadata, otel_trace_carrier, }; - let iter = chunks_contents + let chunks: Vec = chunks_contents .into_iter() .enumerate() .map(move |(chunk_index, uri)| { @@ -470,25 +606,30 @@ pub mod db { Chunk::new(submission_id, chunk_index.try_into().unwrap(), uri) }) .collect(); - insert_submission(submission, iter, &mut conn).await?; - // Empty submissions get special handling: we mark them as completed right away. - // See https://github.com/channable/opsqueue/issues/86 for rationale. - if len == 0 { - match maybe_complete_submission(submission_id, conn).await { - // Forward our database errors to the caller. - Err(E::L(e)) => return Err(e), - // If the submission ID can't be found, that's too bad, but it's not our problem anymore i guess. - Err(E::R(_)) => { - tracing::warn!(%submission_id, "Presumed zero-length submission not found"); - } - // If everything went OK, this *could* still indicate a bug in producer code, so let's just log it. - // Our future selves might thank us. - Ok(true) => { - tracing::debug!(%submission_id, "Zero-length submission marked as completed"); - } - // This should never happen. If it does, better log it. - Ok(false) => { - tracing::warn!(%submission_id, "Zero-length submission wasn't zero-length?!"); + + if paused { + insert_paused_submission(submission, chunks, &mut conn).await?; + } else { + insert_submission(submission, chunks, &mut conn).await?; + // Empty submissions get special handling: we mark them as completed right away. + // See https://github.com/channable/opsqueue/issues/86 for rationale. + if len == 0 { + match maybe_complete_submission(submission_id, conn).await { + // Forward our database errors to the caller. + Err(E::L(e)) => return Err(e), + // If the submission ID can't be found, that's too bad, but it's not our problem anymore i guess. + Err(E::R(_)) => { + tracing::warn!(%submission_id, "Presumed zero-length submission not found"); + } + // If everything went OK, this *could* still indicate a bug in producer code, so let's just log it. + // Our future selves might thank us. + Ok(true) => { + tracing::debug!(%submission_id, "Zero-length submission marked as completed"); + } + // This should never happen. If it does, better log it. + Ok(false) => { + tracing::warn!(%submission_id, "Zero-length submission wasn't zero-length?!"); + } } } } @@ -579,12 +720,15 @@ pub mod db { r#" SELECT id AS "id: SubmissionId" FROM submissions WHERE prefix = $1 UNION ALL - SELECT id AS "id: SubmissionId" FROM submissions_completed WHERE prefix = $2 + SELECT id AS "id: SubmissionId" FROM submissions_paused WHERE prefix = $2 UNION ALL - SELECT id AS "id: SubmissionId" FROM submissions_failed WHERE prefix = $3 + SELECT id AS "id: SubmissionId" FROM submissions_completed WHERE prefix = $3 + UNION ALL + SELECT id AS "id: SubmissionId" FROM submissions_failed WHERE prefix = $4 "#, prefix, prefix, + prefix, prefix ) .fetch_optional(conn.get_inner()) @@ -808,6 +952,40 @@ pub mod db { return Ok(Some(SubmissionStatus::Cancelled(cancelled_submission))); } + let paused_row_opt = query!( + r#" + SELECT + id AS "id: SubmissionId" + , prefix + , chunks_total AS "chunks_total: ChunkCount" + , chunks_done AS "chunks_done: ChunkCount" + , chunk_size AS "chunk_size!: ChunkSize" + , metadata + , ( SELECT json_group_object(metadata_key, metadata_value) + FROM submissions_metadata + WHERE submission_id = submissions_paused.id + ) AS "strategic_metadata!: sqlx::types::Json" + , otel_trace_carrier + FROM submissions_paused WHERE id = $1 + "#, + id + ) + .fetch_optional(conn.get_inner()) + .await?; + if let Some(row) = paused_row_opt { + let paused_submission = SubmissionPaused { + id: row.id, + prefix: row.prefix, + chunks_total: row.chunks_total, + chunks_done: row.chunks_done, + chunk_size: row.chunk_size, + metadata: row.metadata, + strategic_metadata: row.strategic_metadata.0, + otel_trace_carrier: row.otel_trace_carrier, + }; + return Ok(Some(SubmissionStatus::Paused(paused_submission))); + } + Ok(None) } @@ -877,6 +1055,15 @@ 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)), } } @@ -900,6 +1087,22 @@ pub mod db { Ok(()) } + /// Do not call directly! Must be called inside a transaction. + /// + /// # Errors + /// + /// Returns [`DatabaseError`] if any SQL query fails. + /// + /// Returns [`SubmissionNotFound`] if the submission is not found in `submissions_paused`. + pub async fn cancel_paused_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(()) + } + #[tracing::instrument(skip(conn))] pub(super) async fn cancel_submission_raw( id: SubmissionId, @@ -932,6 +1135,38 @@ pub mod db { } } + #[tracing::instrument(skip(conn))] + pub(super) async fn cancel_paused_submission_raw( + id: SubmissionId, + mut conn: impl WriterConnection, + ) -> Result<(), E> { + let now = chrono::prelude::Utc::now(); + + let submission_opt = query!( + " + INSERT INTO submissions_cancelled + (id, chunks_total, prefix, metadata, cancelled_at, chunks_done) + SELECT id, chunks_total, prefix, metadata, julianday($1), chunks_done FROM submissions_paused WHERE id = $2; + + DELETE FROM submissions_paused WHERE id = $3 RETURNING *; + ", + now, + id, + id, + ) + .fetch_optional(conn.get_inner()) + .await?; + if submission_opt.is_none() { + Err(E::R(SubmissionNotFound(id))) + } else { + counter!(crate::prometheus::SUBMISSIONS_CANCELLED_COUNTER).increment(1); + histogram!(crate::prometheus::SUBMISSIONS_DURATION_CANCEL_HISTOGRAM).record( + crate::prometheus::time_delta_as_f64(Utc::now() - id.timestamp()), + ); + Ok(()) + } + } + #[tracing::instrument(skip(conn))] /// Do not call directly! MUST be called inside a transaction. pub(super) async fn complete_submission_raw( @@ -1093,6 +1328,40 @@ pub mod db { Ok(u64::try_from(count).expect("COUNT(*) is always non-negative")) } + /// Count paused submissions. + /// + /// # Errors + /// + /// Returns an error if the count query fails. + /// + /// # Panics + /// + /// Panics if `COUNT(*)` returns a negative value, which `SQLite` never does. + #[tracing::instrument(skip(db))] + pub async fn count_submissions_paused(mut db: impl Connection) -> sqlx::Result { + let count = sqlx::query_scalar!("SELECT COUNT(1) as count FROM submissions_paused;") + .fetch_one(db.get_inner()) + .await?; + Ok(u64::try_from(count).expect("COUNT(*) is always non-negative")) + } + + /// Count cancelled submissions. + /// + /// # Errors + /// + /// Returns an error if the count query fails. + /// + /// # Panics + /// + /// Panics if `COUNT(*)` returns a negative value, which `SQLite` never does. + #[tracing::instrument(skip(db))] + pub async fn count_submissions_cancelled(mut db: impl Connection) -> sqlx::Result { + let count = sqlx::query_scalar!("SELECT COUNT(1) as count FROM submissions_cancelled;") + .fetch_one(db.get_inner()) + .await?; + Ok(u64::try_from(count).expect("COUNT(*) is always non-negative")) + } + /// Transactionally removes all completed/failed submissions, /// including all their chunks and associated strategic metadata. /// @@ -1206,8 +1475,10 @@ pub mod test { use itertools::Itertools; use sqlformat::{FormatOptions, QueryParams, format}; use sqlx::{Row, SqliteConnection}; + 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 super::db::*; @@ -1446,6 +1717,7 @@ pub mod test { None, strategic_metadata.clone(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1521,6 +1793,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1531,6 +1804,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1541,6 +1815,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1551,6 +1826,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1583,6 +1859,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1593,6 +1870,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1603,6 +1881,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1656,6 +1935,7 @@ pub mod test { StrategicMetadataMap::default(), // chunk size ChunkSize::default(), + false, &mut conn, ) .await @@ -1761,4 +2041,125 @@ pub mod test { let deserialized: SubmissionCancelled = serde_json::from_value(json).unwrap(); assert_eq!(deserialized, cancelled); } + + #[sqlx::test(migrator = "crate::MIGRATOR")] + pub async fn test_query_plan_submission_status_paused(db: sqlx::SqlitePool) { + let mut conn = db.acquire().await.unwrap(); + let query = r" + SELECT + id + , prefix + , chunks_total + , chunks_done + , chunk_size + , metadata + , ( SELECT json_group_object(metadata_key, metadata_value) + FROM submissions_metadata + WHERE submission_id = submissions_paused.id + ) AS strategic_metadata + , otel_trace_carrier + FROM submissions_paused WHERE id = 1 + "; + + let explained = explain_query_plan(query, &mut conn).await; + assert_non_regressing_query_plan(query, &explained); + insta::assert_snapshot!(explained, @r" + 3, 0, SEARCH submissions_paused USING INDEX sqlite_autoindex_submissions_paused_1 (id=?) + 15, 0, CORRELATED SCALAR SUBQUERY 1 + 20, 15, SEARCH submissions_metadata USING PRIMARY KEY (submission_id=?) + "); + } + + #[sqlx::test(migrator = "crate::MIGRATOR")] + pub async fn test_unpause_submission(db: sqlx::SqlitePool) { + let db = WriterPool::new(db); + let mut conn = db.writer_conn().await.unwrap(); + let (submission, chunks) = Submission::from_vec( + vec![Some("foo".into()), Some("bar".into()), Some("baz".into())], + None, + ChunkSize::default(), + ) + .unwrap(); + insert_paused_submission(submission.clone(), chunks, &mut conn) + .await + .expect("insertion failed"); + + assert_eq!(count_submissions(&mut conn).await.unwrap(), 0); + assert_eq!(count_submissions_paused(&mut conn).await.unwrap(), 1); + 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(); + 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); + assert_eq!(count_chunks_paused(&mut conn).await.unwrap(), 0); + } + + #[sqlx::test(migrator = "crate::MIGRATOR")] + 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, chunks) = Submission::from_vec( + vec![Some("foo".into()), Some("bar".into()), Some("baz".into())], + None, + ChunkSize::default(), + ) + .unwrap(); + insert_paused_submission(submission.clone(), chunks, &mut conn) + .await + .expect("insertion failed"); + + 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(); + + assert_eq!(count_submissions(&mut conn).await.unwrap(), 0); + assert_eq!(count_submissions_paused(&mut conn).await.unwrap(), 0); + assert_eq!(count_submissions_completed(&mut conn).await.unwrap(), 0); + assert_eq!(count_submissions_failed(&mut conn).await.unwrap(), 0); + assert_eq!(count_submissions_cancelled(&mut conn).await.unwrap(), 1); + assert_eq!(count_chunks(&mut conn).await.unwrap(), 0); + assert_eq!(count_chunks_failed(&mut conn).await.unwrap(), 3); + assert_eq!(count_chunks_paused(&mut conn).await.unwrap(), 0); + } + + #[sqlx::test(migrator = "crate::MIGRATOR")] + pub async fn test_submission_status_paused(db: sqlx::SqlitePool) { + let db = WriterPool::new(db); + let mut conn = db.writer_conn().await.unwrap(); + let (submission, chunks) = Submission::from_vec( + vec![Some("foo".into()), Some("bar".into()), Some("baz".into())], + None, + ChunkSize::default(), + ) + .unwrap(); + insert_paused_submission(submission.clone(), chunks, &mut conn) + .await + .expect("insertion failed"); + + let status = submission_status(submission.id, &mut conn) + .await + .unwrap() + .unwrap(); + assert_matches!(status, SubmissionStatus::Paused(_)); + } + + #[sqlx::test(migrator = "crate::MIGRATOR")] + /// Test that an empty submission inserted in the paused state stays paused + /// (unlike empty non-paused submissions which are auto-completed). + pub async fn insert_empty_paused_submission_stays_paused(db: sqlx::SqlitePool) { + let db = WriterPool::new(db); + let mut conn = db.writer_conn().await.unwrap(); + let (submission, chunks) = + Submission::from_vec(vec![], None, ChunkSize::default()).unwrap(); + insert_paused_submission(submission.clone(), chunks, &mut conn) + .await + .expect("insertion failed"); + + assert_eq!(count_submissions(&mut conn).await.unwrap(), 0); + assert_eq!(count_submissions_paused(&mut conn).await.unwrap(), 1); + assert_eq!(count_submissions_completed(&mut conn).await.unwrap(), 0); + } } diff --git a/opsqueue/src/consumer/client.rs b/opsqueue/src/consumer/client.rs index 94abc811..ac37557c 100644 --- a/opsqueue/src/consumer/client.rs +++ b/opsqueue/src/consumer/client.rs @@ -578,6 +578,7 @@ mod tests { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await diff --git a/opsqueue/src/consumer/strategy.rs b/opsqueue/src/consumer/strategy.rs index a725d138..5ac8ae17 100644 --- a/opsqueue/src/consumer/strategy.rs +++ b/opsqueue/src/consumer/strategy.rs @@ -781,6 +781,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await diff --git a/opsqueue/src/producer/client.rs b/opsqueue/src/producer/client.rs index 1632050b..09e4ab5b 100644 --- a/opsqueue/src/producer/client.rs +++ b/opsqueue/src/producer/client.rs @@ -203,6 +203,49 @@ impl Client { .await } + /// Unpause a paused submission, making it available to consumers again. + /// + /// Returns an error if the submission is not currently paused. + /// + /// # Errors + /// + /// Returns an error if the HTTP request fails or the server returns an unexpected status. + pub async fn unpause_submission( + &self, + submission_id: SubmissionId, + ) -> Result<(), E![SubmissionNotFound, InternalProducerClientError]> { + (|| async { + let base_url = &self.base_url; + let response = self + .http_client + .post(format!("{base_url}/submissions/unpause/{submission_id}")) + .send() + .await + .map_err(|e| R(e.into()))?; + let status = response.status(); + match status { + StatusCode::OK => Ok(()), + StatusCode::NOT_FOUND => { + let not_found_err = response + .json::() + .await + .map_err(|e| R(e.into()))?; + Err(L(not_found_err)) + } + _ => Err(R(InternalProducerClientError::UnexpectedStatus(status))), + } + }) + .retry(retry_policy()) + .when(|e| match e { + L(_) => false, + R(client_err) => client_err.is_ephemeral(), + }) + .notify(|err, dur| { + tracing::debug!("retrying error {err:?} with sleeping {dur:?}"); + }) + .await + } + /// Get the status of an existing submission identified by its `submission_id`. /// /// This uses the GET `/producer/submissions` endpoint. @@ -438,6 +481,7 @@ mod tests { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -467,6 +511,7 @@ mod tests { metadata: None, strategic_metadata: StrategicMetadataMap::default(), chunk_size: None, + paused: false, }; client .insert_submission(&submission, &std::collections::HashMap::default()) @@ -510,6 +555,7 @@ mod tests { metadata: None, strategic_metadata: StrategicMetadataMap::default(), chunk_size: None, + paused: false, }; let submission_id = client .insert_submission(&submission, &std::collections::HashMap::default()) @@ -524,7 +570,8 @@ mod tests { match status { SubmissionStatus::Completed(_) | SubmissionStatus::Failed(_, _) - | SubmissionStatus::Cancelled(_) => { + | SubmissionStatus::Cancelled(_) + | SubmissionStatus::Paused(_) => { panic!("Expected a SubmissionStatus that is still Inprogress, got: {status:?}"); } SubmissionStatus::InProgress(submission) => { @@ -534,4 +581,80 @@ mod tests { } } } + + #[sqlx::test(migrator = "crate::MIGRATOR")] + async fn test_insert_paused_submission_and_unpause(pool: sqlx::SqlitePool) { + let url = "0.0.0.0:4003"; + start_server_in_background(&pool, url).await; + let client = Client::new(url); + + let pool = WriterPool::new(pool); + let mut conn = pool.writer_conn().await.unwrap(); + let count = submission::db::count_submissions(&mut conn) + .await + .expect("Should be OK"); + assert_eq!(count, 0); + + let submission = InsertSubmission { + chunk_contents: ChunkContents::Direct { + contents: vec![None, None, None], + }, + metadata: None, + strategic_metadata: StrategicMetadataMap::default(), + chunk_size: None, + paused: true, + }; + let submission_id = client + .insert_submission(&submission, &std::collections::HashMap::default()) + .await + .expect("Should be OK"); + + let count = submission::db::count_submissions_paused(&mut conn) + .await + .expect("Should be OK"); + assert_eq!(count, 1); + + let status: SubmissionStatus = client + .get_submission(submission_id) + .await + .expect("Should be OK") + .expect("Should be Some"); + match status { + SubmissionStatus::Completed(_) + | SubmissionStatus::Failed(_, _) + | SubmissionStatus::Cancelled(_) + | SubmissionStatus::InProgress(_) => { + panic!("Expected a SubmissionStatus that is Paused, got: {status:?}"); + } + SubmissionStatus::Paused(submission) => { + assert_eq!(submission.chunks_done, 0); + assert_eq!(submission.chunks_total, 3); + assert_eq!(submission.id, submission_id); + } + } + + client + .unpause_submission(submission_id) + .await + .expect("Should be OK"); + + let status: SubmissionStatus = client + .get_submission(submission_id) + .await + .expect("Should be OK") + .expect("Should be Some"); + match status { + SubmissionStatus::Completed(_) + | SubmissionStatus::Failed(_, _) + | SubmissionStatus::Cancelled(_) + | SubmissionStatus::Paused(_) => { + panic!("Expected a SubmissionStatus that is InProgress, got: {status:?}"); + } + SubmissionStatus::InProgress(submission) => { + assert_eq!(submission.chunks_done, 0); + assert_eq!(submission.chunks_total, 3); + assert_eq!(submission.id, submission_id); + } + } + } } diff --git a/opsqueue/src/producer/common.rs b/opsqueue/src/producer/common.rs index 41f3c8e9..f98b4b86 100644 --- a/opsqueue/src/producer/common.rs +++ b/opsqueue/src/producer/common.rs @@ -10,6 +10,10 @@ pub struct InsertSubmission { #[serde(default)] pub strategic_metadata: StrategicMetadataMap, pub chunk_size: Option, + /// When `true`, the submission is inserted in a paused state and will not + /// be dispatched to consumers until explicitly unpaused. + #[serde(default)] + pub paused: bool, } /// Either embedded chunk contents or a reference to object storage. diff --git a/opsqueue/src/producer/server.rs b/opsqueue/src/producer/server.rs index 1fa4d678..7e603233 100644 --- a/opsqueue/src/producer/server.rs +++ b/opsqueue/src/producer/server.rs @@ -67,6 +67,10 @@ impl ServerState { "/submissions/cancel/{submission_id}", post(cancel_submission), ) + .route( + "/submissions/unpause/{submission_id}", + post(unpause_submission), + ) .route( "/submissions/count_completed", get(submissions_count_completed), @@ -138,6 +142,29 @@ async fn cancel_submission( } } +/// 200 if the submission was successfully unpaused. +/// 404 if the submission could not be found in the paused state. +/// 500 if a `DatabaseError` occurred. +async fn unpause_submission( + State(state): State, + Path(submission_id): Path, +) -> Result<(), Response> { + let mut conn = state + .pool + .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()), + } +} + async fn submission_status( State(state): State, Path(submission_id): Path, @@ -199,6 +226,7 @@ async fn insert_submission( request.metadata, request.strategic_metadata, request.chunk_size.unwrap_or_default(), + request.paused, &mut conn, ) .await?; @@ -208,8 +236,10 @@ async fn insert_submission( // this is the moment to perform an extra WAL checkpoint let _ = db::perform_explicit_wal_checkpoint(conn).await; - // We've done a new insert! Let's tell any waiting consumers! - state.notify_on_insert.notify_waiters(); + // Notify waiting consumers, but only for non-paused submissions. + if !request.paused { + state.notify_on_insert.notify_waiters(); + } Ok(Json(submission_id)) } diff --git a/opsqueue/src/prometheus.rs b/opsqueue/src/prometheus.rs index d4b308a1..29f63baf 100644 --- a/opsqueue/src/prometheus.rs +++ b/opsqueue/src/prometheus.rs @@ -19,6 +19,8 @@ pub const SUBMISSIONS_TOTAL_COUNTER: &str = "submissions_total_count"; pub const SUBMISSIONS_COMPLETED_COUNTER: &str = "submissions_completed_count"; pub const SUBMISSIONS_FAILED_COUNTER: &str = "submissions_failed_count"; pub const SUBMISSIONS_CANCELLED_COUNTER: &str = "submissions_cancelled_count"; +pub const SUBMISSIONS_PAUSED_COUNTER: &str = "submissions_paused_count"; +pub const SUBMISSIONS_UNPAUSED_COUNTER: &str = "submissions_unpaused_count"; pub const SUBMISSIONS_DURATION_COMPLETE_HISTOGRAM: &str = "submissions_complete_duration_seconds"; pub const SUBMISSIONS_DURATION_FAIL_HISTOGRAM: &str = "submissions_fail_duration_seconds"; pub const SUBMISSIONS_DURATION_CANCEL_HISTOGRAM: &str = "submissions_cancel_duration_seconds"; @@ -67,6 +69,16 @@ pub fn describe_metrics() { Unit::Count, "Number of submissions cancelled (client-requested cancellation, not failure) permanently" ); + describe_counter!( + SUBMISSIONS_PAUSED_COUNTER, + Unit::Count, + "Number of submissions paused" + ); + describe_counter!( + SUBMISSIONS_UNPAUSED_COUNTER, + Unit::Count, + "Number of submissions unpaused (resumed)" + ); describe_histogram!( SUBMISSIONS_DURATION_COMPLETE_HISTOGRAM, Unit::Seconds, From ad14db989fee9e93fd41e16540e2088663715c8c Mon Sep 17 00:00:00 2001 From: Sem Mulder Date: Wed, 5 Aug 2026 12:06:01 +0200 Subject: [PATCH 05/11] Address feedback --- libs/opsqueue_python/src/producer.rs | 2 +- opsqueue/src/common/chunk.rs | 4 +-- opsqueue/src/common/submission.rs | 41 +++++++++++++--------------- opsqueue/src/producer/client.rs | 2 +- opsqueue/src/prometheus.rs | 2 +- 5 files changed, 24 insertions(+), 27 deletions(-) diff --git a/libs/opsqueue_python/src/producer.rs b/libs/opsqueue_python/src/producer.rs index 1c40ad25..fcf2b637 100644 --- a/libs/opsqueue_python/src/producer.rs +++ b/libs/opsqueue_python/src/producer.rs @@ -159,7 +159,7 @@ impl ProducerClient { }) } - /// Unpause a paused submission, making it available to consumers again. + /// Unpause a paused submission, making it available to consumers. /// /// Will return an error if the submission is not currently paused. /// diff --git a/opsqueue/src/common/chunk.rs b/opsqueue/src/common/chunk.rs index 2526f5aa..9fa4e88b 100644 --- a/opsqueue/src/common/chunk.rs +++ b/opsqueue/src/common/chunk.rs @@ -673,7 +673,7 @@ pub mod db { " INSERT INTO chunks_failed (submission_id, chunk_index, input_content, failure, skipped, failed_at) - SELECT submission_id, chunk_index, input_content, '', 1, julianday($1) FROM chunks_paused WHERE submission_id = $2; + SELECT submission_id, chunk_index, input_content, '', TRUE, julianday($1) FROM chunks_paused WHERE submission_id = $2; DELETE FROM chunks_paused WHERE submission_id = $3; ", @@ -705,7 +705,7 @@ pub mod db { INSERT INTO chunks_failed (submission_id, chunk_index, input_content, failure, skipped, failed_at) - SELECT submission_id, chunk_index, input_content, '', 1, julianday($1) FROM chunks WHERE chunks.submission_id = $2; + SELECT submission_id, chunk_index, input_content, '', TRUE, julianday($1) FROM chunks WHERE chunks.submission_id = $2; DELETE FROM chunks WHERE chunks.submission_id = $3; diff --git a/opsqueue/src/common/submission.rs b/opsqueue/src/common/submission.rs index 82fa39a9..40746d64 100644 --- a/opsqueue/src/common/submission.rs +++ b/opsqueue/src/common/submission.rs @@ -537,21 +537,21 @@ pub mod db { id: SubmissionId, mut conn: impl WriterConnection, ) -> Result<(), E> { - let row = query!( + let res = query!( " INSERT INTO submissions (id, chunks_total, chunks_done, prefix, metadata, otel_trace_carrier, chunk_size) SELECT id, chunks_total, chunks_done, prefix, metadata, otel_trace_carrier, chunk_size FROM submissions_paused WHERE id = $1; - DELETE FROM submissions_paused WHERE id = $2 RETURNING *; + DELETE FROM submissions_paused WHERE id = $2; ", id, id, ) - .fetch_optional(conn.get_inner()) + .execute(conn.get_inner()) .await?; - if row.is_none() { + if res.rows_affected() == 0 { Err(E::R(SubmissionNotFound(id))) } else { counter!(crate::prometheus::SUBMISSIONS_UNPAUSED_COUNTER).increment(1); @@ -617,7 +617,7 @@ pub mod db { match maybe_complete_submission(submission_id, conn).await { // Forward our database errors to the caller. Err(E::L(e)) => return Err(e), - // If the submission ID can't be found, that's too bad, but it's not our problem anymore i guess. + // If the submission ID can't be found, that's too bad, but it's not our problem anymore I guess. Err(E::R(_)) => { tracing::warn!(%submission_id, "Presumed zero-length submission not found"); } @@ -720,16 +720,13 @@ pub mod db { r#" SELECT id AS "id: SubmissionId" FROM submissions WHERE prefix = $1 UNION ALL - SELECT id AS "id: SubmissionId" FROM submissions_paused WHERE prefix = $2 + SELECT id AS "id: SubmissionId" FROM submissions_paused WHERE prefix = $1 UNION ALL - SELECT id AS "id: SubmissionId" FROM submissions_completed WHERE prefix = $3 + SELECT id AS "id: SubmissionId" FROM submissions_completed WHERE prefix = $1 UNION ALL - SELECT id AS "id: SubmissionId" FROM submissions_failed WHERE prefix = $4 + SELECT id AS "id: SubmissionId" FROM submissions_failed WHERE prefix = $1 "#, prefix, - prefix, - prefix, - prefix ) .fetch_optional(conn.get_inner()) .await?; @@ -1078,7 +1075,7 @@ pub mod db { /// # Errors /// /// Returns an error if cancellation or chunk skipping fails. - pub async fn cancel_submission_notx( + async fn cancel_submission_notx( id: SubmissionId, mut conn: impl WriterConnection, ) -> Result<(), E> { @@ -1094,7 +1091,7 @@ pub mod db { /// Returns [`DatabaseError`] if any SQL query fails. /// /// Returns [`SubmissionNotFound`] if the submission is not found in `submissions_paused`. - pub async fn cancel_paused_submission_notx( + async fn cancel_paused_submission_notx( id: SubmissionId, mut conn: impl WriterConnection, ) -> Result<(), E> { @@ -1110,21 +1107,21 @@ pub mod db { ) -> Result<(), E> { let now = chrono::prelude::Utc::now(); - let submission_opt = query!( + let res = query!( " INSERT INTO submissions_cancelled (id, chunks_total, prefix, metadata, cancelled_at, chunks_done) SELECT id, chunks_total, prefix, metadata, julianday($1), chunks_done FROM submissions WHERE id = $2; - DELETE FROM submissions WHERE id = $3 RETURNING *; + DELETE FROM submissions WHERE id = $3; ", now, id, id, ) - .fetch_optional(conn.get_inner()) + .execute(conn.get_inner()) .await?; - if submission_opt.is_none() { + if res.rows_affected() == 0 { Err(E::R(SubmissionNotFound(id))) } else { counter!(crate::prometheus::SUBMISSIONS_CANCELLED_COUNTER).increment(1); @@ -1142,21 +1139,21 @@ pub mod db { ) -> Result<(), E> { let now = chrono::prelude::Utc::now(); - let submission_opt = query!( + let res = query!( " INSERT INTO submissions_cancelled (id, chunks_total, prefix, metadata, cancelled_at, chunks_done) SELECT id, chunks_total, prefix, metadata, julianday($1), chunks_done FROM submissions_paused WHERE id = $2; - DELETE FROM submissions_paused WHERE id = $3 RETURNING *; + DELETE FROM submissions_paused WHERE id = $3; ", now, id, id, ) - .fetch_optional(conn.get_inner()) + .execute(conn.get_inner()) .await?; - if submission_opt.is_none() { + if res.rows_affected() == 0 { Err(E::R(SubmissionNotFound(id))) } else { counter!(crate::prometheus::SUBMISSIONS_CANCELLED_COUNTER).increment(1); @@ -1260,7 +1257,7 @@ pub mod db { /// # Errors /// /// Returns an error if submission/chunk failure transitions cannot be persisted. - pub async fn fail_submission_notx( + pub(crate) async fn fail_submission_notx( id: SubmissionId, failed_chunk_index: ChunkIndex, failure: String, diff --git a/opsqueue/src/producer/client.rs b/opsqueue/src/producer/client.rs index 09e4ab5b..17a94691 100644 --- a/opsqueue/src/producer/client.rs +++ b/opsqueue/src/producer/client.rs @@ -203,7 +203,7 @@ impl Client { .await } - /// Unpause a paused submission, making it available to consumers again. + /// Unpause a paused submission, making it available to consumers. /// /// Returns an error if the submission is not currently paused. /// diff --git a/opsqueue/src/prometheus.rs b/opsqueue/src/prometheus.rs index 29f63baf..99e76439 100644 --- a/opsqueue/src/prometheus.rs +++ b/opsqueue/src/prometheus.rs @@ -77,7 +77,7 @@ pub fn describe_metrics() { describe_counter!( SUBMISSIONS_UNPAUSED_COUNTER, Unit::Count, - "Number of submissions unpaused (resumed)" + "Number of submissions unpaused" ); describe_histogram!( SUBMISSIONS_DURATION_COMPLETE_HISTOGRAM, From 22c5d1edaa6d8df40f9f2bd46c9aa3ec6175af91 Mon Sep 17 00:00:00 2001 From: Sem Mulder Date: Wed, 5 Aug 2026 13:06:38 +0200 Subject: [PATCH 06/11] Pull out queries and EXPLAIN the now shared query --- opsqueue/src/common/submission.rs | 487 ++++++++++++++++-------------- 1 file changed, 259 insertions(+), 228 deletions(-) diff --git a/opsqueue/src/common/submission.rs b/opsqueue/src/common/submission.rs index 40746d64..fab26829 100644 --- a/opsqueue/src/common/submission.rs +++ b/opsqueue/src/common/submission.rs @@ -310,7 +310,7 @@ pub mod db { }; use axum_prometheus::metrics::{counter, histogram}; use chunk::ChunkSize; - use sqlx::{QueryBuilder, Sqlite, query, query_scalar}; + use sqlx::{Database, QueryBuilder, Sqlite, query, query_as, query_scalar}; use super::{ Chunk, ChunkCount, ChunkIndex, DateTime, Duration, E, Metadata, Submission, @@ -806,7 +806,122 @@ pub mod db { // NOTE: The order is important here; a concurrent writer could move a submission // from InProgress to Completed/Failed in-between the queries. - let submission_row = query!( + let submission_row = submission_status_in_progress_query(id) + .fetch_optional(conn.get_inner()) + .await?; + if let Some(row) = submission_row { + let submission = Submission { + id: row.id, + prefix: row.prefix, + chunks_total: row.chunks_total, + chunks_done: row.chunks_done, + chunk_size: row.chunk_size, + metadata: row.metadata, + strategic_metadata: row.strategic_metadata.0, + otel_trace_carrier: row.otel_trace_carrier, + }; + return Ok(Some(SubmissionStatus::InProgress(submission))); + } + + let completed_row_opt = submission_status_completed_query(id) + .fetch_optional(conn.get_inner()) + .await?; + if let Some(row) = completed_row_opt { + let submission_completed = SubmissionCompleted { + id: row.id, + prefix: row.prefix, + chunks_total: row.chunks_total, + chunk_size: row.chunk_size, + metadata: row.metadata, + strategic_metadata: row.strategic_metadata.0, + completed_at: row.completed_at, + otel_trace_carrier: row.otel_trace_carrier, + }; + return Ok(Some(SubmissionStatus::Completed(submission_completed))); + } + + let failed_row_opt = submission_status_failed_query(id) + .fetch_optional(conn.get_inner()) + .await?; + if let Some(row) = failed_row_opt { + let failed_submission = SubmissionFailed { + id: row.id, + prefix: row.prefix, + chunks_total: row.chunks_total, + chunks_done: row.chunks_done, + chunk_size: row.chunk_size, + metadata: row.metadata, + strategic_metadata: row.strategic_metadata.0, + failed_at: row.failed_at, + failed_chunk_id: row.failed_chunk_id, + otel_trace_carrier: row.otel_trace_carrier, + }; + let failed_chunk_id = (row.id, row.failed_chunk_id).into(); + let failed_chunk = super::chunk::db::get_chunk_failed(failed_chunk_id, conn).await?; + return Ok(Some(SubmissionStatus::Failed( + failed_submission, + failed_chunk, + ))); + } + + let cancelled_row_opt = submission_status_cancelled_query(id) + .fetch_optional(conn.get_inner()) + .await?; + if let Some(row) = cancelled_row_opt { + let cancelled_submission = SubmissionCancelled { + id: row.id, + prefix: row.prefix, + chunks_total: row.chunks_total, + chunks_done: row.chunks_done, + metadata: row.metadata, + strategic_metadata: row.strategic_metadata.0, + cancelled_at: row.cancelled_at, + }; + return Ok(Some(SubmissionStatus::Cancelled(cancelled_submission))); + } + + let paused_row_opt = submission_status_paused_query(id) + .fetch_optional(conn.get_inner()) + .await?; + if let Some(row) = paused_row_opt { + let paused_submission = SubmissionPaused { + id: row.id, + prefix: row.prefix, + chunks_total: row.chunks_total, + chunks_done: row.chunks_done, + chunk_size: row.chunk_size, + metadata: row.metadata, + strategic_metadata: row.strategic_metadata.0, + otel_trace_carrier: row.otel_trace_carrier, + }; + return Ok(Some(SubmissionStatus::Paused(paused_submission))); + } + + Ok(None) + } + + pub(crate) struct SubmissionStatusInProgressRow { + id: SubmissionId, + prefix: Option, + chunks_total: ChunkCount, + chunks_done: ChunkCount, + chunk_size: ChunkSize, + metadata: Option, + strategic_metadata: sqlx::types::Json, + otel_trace_carrier: String, + } + + #[allow(clippy::type_complexity)] + pub(crate) fn submission_status_in_progress_query( + id: SubmissionId, + ) -> query::Map< + 'static, + Sqlite, + fn(::Row) -> Result, + ::Arguments, + > { + query_as!( + SubmissionStatusInProgressRow, r#" SELECT id AS "id: SubmissionId" @@ -824,23 +939,30 @@ pub mod db { "#, id ) - .fetch_optional(conn.get_inner()) - .await?; - if let Some(row) = submission_row { - let submission = Submission { - id: row.id, - prefix: row.prefix, - chunks_total: row.chunks_total, - chunks_done: row.chunks_done, - chunk_size: row.chunk_size, - metadata: row.metadata, - strategic_metadata: row.strategic_metadata.0, - otel_trace_carrier: row.otel_trace_carrier, - }; - return Ok(Some(SubmissionStatus::InProgress(submission))); - } + } - let completed_row_opt = query!( + pub(crate) struct SubmissionStatusCompletedRow { + id: SubmissionId, + prefix: Option, + chunks_total: ChunkCount, + chunk_size: ChunkSize, + metadata: Option, + strategic_metadata: sqlx::types::Json, + completed_at: DateTime, + otel_trace_carrier: String, + } + + #[allow(clippy::type_complexity)] + pub(crate) fn submission_status_completed_query( + id: SubmissionId, + ) -> query::Map< + 'static, + Sqlite, + fn(::Row) -> Result, + ::Arguments, + > { + query_as!( + SubmissionStatusCompletedRow, r#" SELECT id AS "id: SubmissionId" @@ -858,23 +980,32 @@ pub mod db { "#, id ) - .fetch_optional(conn.get_inner()) - .await?; - if let Some(row) = completed_row_opt { - let submission_completed = SubmissionCompleted { - id: row.id, - prefix: row.prefix, - chunks_total: row.chunks_total, - chunk_size: row.chunk_size, - metadata: row.metadata, - strategic_metadata: row.strategic_metadata.0, - completed_at: row.completed_at, - otel_trace_carrier: row.otel_trace_carrier, - }; - return Ok(Some(SubmissionStatus::Completed(submission_completed))); - } + } - let failed_row_opt = query!( + pub(crate) struct SubmissionStatusFailedRow { + id: SubmissionId, + prefix: Option, + chunks_total: ChunkCount, + chunks_done: Option, + chunk_size: ChunkSize, + metadata: Option, + strategic_metadata: sqlx::types::Json, + failed_at: DateTime, + failed_chunk_id: ChunkIndex, + otel_trace_carrier: String, + } + + #[allow(clippy::type_complexity)] + pub(crate) fn submission_status_failed_query( + id: SubmissionId, + ) -> query::Map< + 'static, + Sqlite, + fn(::Row) -> Result, + ::Arguments, + > { + query_as!( + SubmissionStatusFailedRow, r#" SELECT id AS "id: SubmissionId" @@ -894,30 +1025,29 @@ pub mod db { "#, id ) - .fetch_optional(conn.get_inner()) - .await?; - if let Some(row) = failed_row_opt { - let failed_submission = SubmissionFailed { - id: row.id, - prefix: row.prefix, - chunks_total: row.chunks_total, - chunks_done: row.chunks_done, - chunk_size: row.chunk_size, - metadata: row.metadata, - strategic_metadata: row.strategic_metadata.0, - failed_at: row.failed_at, - failed_chunk_id: row.failed_chunk_id, - otel_trace_carrier: row.otel_trace_carrier, - }; - let failed_chunk_id = (row.id, row.failed_chunk_id).into(); - let failed_chunk = super::chunk::db::get_chunk_failed(failed_chunk_id, conn).await?; - return Ok(Some(SubmissionStatus::Failed( - failed_submission, - failed_chunk, - ))); - } + } + + pub(crate) struct SubmissionStatusCancelledRow { + id: SubmissionId, + prefix: Option, + chunks_total: ChunkCount, + chunks_done: ChunkCount, + metadata: Option, + strategic_metadata: sqlx::types::Json, + cancelled_at: DateTime, + } - let cancelled_row_opt = query!( + #[allow(clippy::type_complexity)] + pub(crate) fn submission_status_cancelled_query( + id: SubmissionId, + ) -> query::Map< + 'static, + Sqlite, + fn(::Row) -> Result, + ::Arguments, + > { + query_as!( + SubmissionStatusCancelledRow, r#" SELECT id AS "id: SubmissionId" @@ -934,22 +1064,30 @@ pub mod db { "#, id ) - .fetch_optional(conn.get_inner()) - .await?; - if let Some(row) = cancelled_row_opt { - let cancelled_submission = SubmissionCancelled { - id: row.id, - prefix: row.prefix, - chunks_total: row.chunks_total, - chunks_done: row.chunks_done, - metadata: row.metadata, - strategic_metadata: row.strategic_metadata.0, - cancelled_at: row.cancelled_at, - }; - return Ok(Some(SubmissionStatus::Cancelled(cancelled_submission))); - } + } - let paused_row_opt = query!( + pub(crate) struct SubmissionStatusPausedRow { + id: SubmissionId, + prefix: Option, + chunks_total: ChunkCount, + chunks_done: ChunkCount, + chunk_size: ChunkSize, + metadata: Option, + strategic_metadata: sqlx::types::Json, + otel_trace_carrier: String, + } + + #[allow(clippy::type_complexity)] + pub(crate) fn submission_status_paused_query( + id: SubmissionId, + ) -> query::Map< + 'static, + Sqlite, + fn(::Row) -> Result, + ::Arguments, + > { + query_as!( + SubmissionStatusPausedRow, r#" SELECT id AS "id: SubmissionId" @@ -967,23 +1105,6 @@ pub mod db { "#, id ) - .fetch_optional(conn.get_inner()) - .await?; - if let Some(row) = paused_row_opt { - let paused_submission = SubmissionPaused { - id: row.id, - prefix: row.prefix, - chunks_total: row.chunks_total, - chunks_done: row.chunks_done, - chunk_size: row.chunk_size, - metadata: row.metadata, - strategic_metadata: row.strategic_metadata.0, - otel_trace_carrier: row.otel_trace_carrier, - }; - return Ok(Some(SubmissionStatus::Paused(paused_submission))); - } - - Ok(None) } #[tracing::instrument(skip(conn))] @@ -1471,7 +1592,7 @@ pub mod test { use chunk::ChunkSize; use itertools::Itertools; use sqlformat::{FormatOptions, QueryParams, format}; - use sqlx::{Row, SqliteConnection}; + use sqlx::{Execute, Row, Sqlite}; use std::assert_matches; use crate::common::StrategicMetadataMap; @@ -1481,40 +1602,36 @@ pub mod test { use super::db::*; use super::*; - async fn explain_query_plan(query: &str, conn: &mut SqliteConnection) -> String { - sqlx::raw_sql(sqlx::AssertSqlSafe(format!("EXPLAIN QUERY PLAN {query}"))) - .fetch_all(&mut *conn) - .await - .unwrap_or_else(|_| panic!("Invalid query: \n{query}\n")) - .into_iter() - .map(|row| { - let id = row.get::("id"); - let parent = row.get::("parent"); - let detail = row.get::("detail"); - format!("{id}, {parent}, {detail}") - }) - .join("\n") - } - - fn assert_non_regressing_query_plan(query: &str, explained: &str) { - assert!( - !explained.contains("MATERIALIZED"), - "Query should contain no materialization, but it did.\n\nQuery: {query}\n\nPlan:\n\n{explained}" - ); - assert!( - !explained.contains("B-TREE"), - "Query should contain no temporary B-tree construction, but it did.\n\nQuery: {query}\n\nPlan:\n\n{explained}" - ); + async fn explain_query_plan<'q, Q: Execute<'q, Sqlite>>( + query: Q, + db: sqlx::SqlitePool, + ) -> String { + let mut conn = db.acquire().await.unwrap(); + let query = query.sql(); + let query_string = query.as_str(); + sqlx::raw_sql(sqlx::AssertSqlSafe(format!( + "EXPLAIN QUERY PLAN {query_string}" + ))) + .fetch_all(&mut *conn) + .await + .unwrap_or_else(|_| panic!("Invalid query: \n{query_string}\n")) + .into_iter() + .map(|row| { + let id = row.get::("id"); + let parent = row.get::("parent"); + let detail = row.get::("detail"); + format!("{id}, {parent}, {detail}") + }) + .join("\n") } #[sqlx::test(migrator = "crate::MIGRATOR")] pub async fn test_query_plan_lookup_by_strategic_metadata(db: sqlx::SqlitePool) { - let mut conn = db.acquire().await.unwrap(); let strategic_metadata: StrategicMetadataMap = [("company_id".to_string(), 1), ("project_id".to_string(), 2)] .into_iter() .collect(); - let qb = lookup_ids_by_strategic_metadata_query(&strategic_metadata, 100_000); + let mut qb = lookup_ids_by_strategic_metadata_query(&strategic_metadata, 100_000); let options = FormatOptions::default(); let formatted_query = format(qb.sql().as_str(), &QueryParams::None, &options); insta::assert_snapshot!(formatted_query, @" @@ -1533,8 +1650,7 @@ pub mod test { LIMIT ? "); - let explained = explain_query_plan(&formatted_query, &mut conn).await; - assert_non_regressing_query_plan(&formatted_query, &explained); + let explained = explain_query_plan(qb.build_query_scalar::(), db).await; insta::assert_snapshot!(explained, @" 8, 0, SEARCH s0 USING COVERING INDEX lookup_submission_by_metadata (metadata_key=? AND metadata_value=?) 16, 0, SEARCH submissions USING COVERING INDEX sqlite_autoindex_submissions_1 (id=?) @@ -1544,114 +1660,46 @@ pub mod test { #[sqlx::test(migrator = "crate::MIGRATOR")] pub async fn test_query_plan_submission_status_in_progress(db: sqlx::SqlitePool) { - let mut conn = db.acquire().await.unwrap(); - let query = r" - SELECT - id - , prefix - , chunks_total - , chunks_done - , chunk_size - , metadata - , ( SELECT json_group_object(metadata_key, metadata_value) - FROM submissions_metadata - WHERE submission_id = submissions.id - ) AS strategic_metadata - , otel_trace_carrier - FROM submissions WHERE id = 1 - "; - - let explained = explain_query_plan(query, &mut conn).await; - assert_non_regressing_query_plan(query, &explained); - insta::assert_snapshot!(explained, @r" + let query = submission_status_in_progress_query(SubmissionId::new()); + let explained = explain_query_plan(query, db).await; + insta::assert_snapshot!(explained, @" 3, 0, SEARCH submissions USING INDEX sqlite_autoindex_submissions_1 (id=?) - 15, 0, CORRELATED SCALAR SUBQUERY 1 - 20, 15, SEARCH submissions_metadata USING PRIMARY KEY (submission_id=?) + 17, 0, CORRELATED SCALAR SUBQUERY 1 + 22, 17, SEARCH submissions_metadata USING PRIMARY KEY (submission_id=?) "); } #[sqlx::test(migrator = "crate::MIGRATOR")] pub async fn test_query_plan_submission_status_completed(db: sqlx::SqlitePool) { - let mut conn = db.acquire().await.unwrap(); - let query = r" - SELECT - id - , prefix - , chunks_total - , chunk_size - , metadata - , ( SELECT json_group_object(metadata_key, metadata_value) - FROM submissions_metadata - WHERE submission_id = submissions_completed.id - ) AS strategic_metadata - , completed_at - , otel_trace_carrier - FROM submissions_completed WHERE id = 1 - "; + let query = submission_status_completed_query(SubmissionId::new()); - let explained = explain_query_plan(query, &mut conn).await; - assert_non_regressing_query_plan(query, &explained); - insta::assert_snapshot!(explained, @r" + let explained = explain_query_plan(query, db).await; + insta::assert_snapshot!(explained, @" 3, 0, SEARCH submissions_completed USING INDEX sqlite_autoindex_submissions_completed_1 (id=?) - 14, 0, CORRELATED SCALAR SUBQUERY 1 - 19, 14, SEARCH submissions_metadata USING PRIMARY KEY (submission_id=?) + 16, 0, CORRELATED SCALAR SUBQUERY 1 + 21, 16, SEARCH submissions_metadata USING PRIMARY KEY (submission_id=?) "); } #[sqlx::test(migrator = "crate::MIGRATOR")] pub async fn test_query_plan_submission_status_failed(db: sqlx::SqlitePool) { - let mut conn = db.acquire().await.unwrap(); - let query = r" - SELECT - id - , prefix - , chunks_total - , chunks_done - , chunk_size - , metadata - , ( SELECT json_group_object(metadata_key, metadata_value) - FROM submissions_metadata - WHERE submission_id = submissions_failed.id - ) AS strategic_metadata - , failed_at - , failed_chunk_id - , otel_trace_carrier - FROM submissions_failed WHERE id = 1 - "; - - let explained = explain_query_plan(query, &mut conn).await; - assert_non_regressing_query_plan(query, &explained); - insta::assert_snapshot!(explained, @r" + let query = submission_status_failed_query(SubmissionId::new()); + let explained = explain_query_plan(query, db).await; + insta::assert_snapshot!(explained, @" 3, 0, SEARCH submissions_failed USING INDEX sqlite_autoindex_submissions_failed_1 (id=?) - 15, 0, CORRELATED SCALAR SUBQUERY 1 - 20, 15, SEARCH submissions_metadata USING PRIMARY KEY (submission_id=?) + 17, 0, CORRELATED SCALAR SUBQUERY 1 + 22, 17, SEARCH submissions_metadata USING PRIMARY KEY (submission_id=?) "); } #[sqlx::test(migrator = "crate::MIGRATOR")] pub async fn test_query_plan_submission_status_cancelled(db: sqlx::SqlitePool) { - let mut conn = db.acquire().await.unwrap(); - let query = r" - SELECT - id - , prefix - , chunks_total - , chunks_done - , metadata - , ( SELECT json_group_object(metadata_key, metadata_value) - FROM submissions_metadata - WHERE submission_id = submissions_cancelled.id - ) AS strategic_metadata - , cancelled_at - FROM submissions_cancelled WHERE id = 1 - "; - - let explained = explain_query_plan(query, &mut conn).await; - assert_non_regressing_query_plan(query, &explained); - insta::assert_snapshot!(explained, @r" + let query = submission_status_cancelled_query(SubmissionId::new()); + let explained = explain_query_plan(query, db).await; + insta::assert_snapshot!(explained, @" 3, 0, SEARCH submissions_cancelled USING INDEX sqlite_autoindex_submissions_cancelled_1 (id=?) - 14, 0, CORRELATED SCALAR SUBQUERY 1 - 19, 14, SEARCH submissions_metadata USING PRIMARY KEY (submission_id=?) + 16, 0, CORRELATED SCALAR SUBQUERY 1 + 21, 16, SEARCH submissions_metadata USING PRIMARY KEY (submission_id=?) "); } @@ -2041,29 +2089,12 @@ pub mod test { #[sqlx::test(migrator = "crate::MIGRATOR")] pub async fn test_query_plan_submission_status_paused(db: sqlx::SqlitePool) { - let mut conn = db.acquire().await.unwrap(); - let query = r" - SELECT - id - , prefix - , chunks_total - , chunks_done - , chunk_size - , metadata - , ( SELECT json_group_object(metadata_key, metadata_value) - FROM submissions_metadata - WHERE submission_id = submissions_paused.id - ) AS strategic_metadata - , otel_trace_carrier - FROM submissions_paused WHERE id = 1 - "; - - let explained = explain_query_plan(query, &mut conn).await; - assert_non_regressing_query_plan(query, &explained); - insta::assert_snapshot!(explained, @r" + let query = submission_status_paused_query(SubmissionId::new()); + let explained = explain_query_plan(query, db).await; + insta::assert_snapshot!(explained, @" 3, 0, SEARCH submissions_paused USING INDEX sqlite_autoindex_submissions_paused_1 (id=?) - 15, 0, CORRELATED SCALAR SUBQUERY 1 - 20, 15, SEARCH submissions_metadata USING PRIMARY KEY (submission_id=?) + 17, 0, CORRELATED SCALAR SUBQUERY 1 + 22, 17, SEARCH submissions_metadata USING PRIMARY KEY (submission_id=?) "); } From 3ee2454b371d0f4121d6f35d13c4ab39f985c031 Mon Sep 17 00:00:00 2001 From: Sem Mulder Date: Wed, 5 Aug 2026 15:26:08 +0200 Subject: [PATCH 07/11] Address Copilot's suppressed comments --- libs/opsqueue_python/src/errors.rs | 9 +- libs/opsqueue_python/src/producer.rs | 27 +++-- libs/opsqueue_python/tests/test_roundtrip.py | 7 +- opsqueue/src/common/chunk.rs | 120 +++++++++++++++---- opsqueue/src/common/submission.rs | 63 +++++++--- 5 files changed, 169 insertions(+), 57 deletions(-) diff --git a/libs/opsqueue_python/src/errors.rs b/libs/opsqueue_python/src/errors.rs index b0de8e7e..043ac11f 100644 --- a/libs/opsqueue_python/src/errors.rs +++ b/libs/opsqueue_python/src/errors.rs @@ -1,12 +1,13 @@ /// NOTE: We define the potentially raisable errors/exceptions in Python /// so we have nice IDE support for docs-on-hover and for 'go to definition'. use std::error::Error; +use std::time::TryFromFloatSecsError; use opsqueue::common::errors::{ E, IncorrectUsage, SubmissionNotCancellable, SubmissionNotFound, TooManyMatchingSubmissions, UnexpectedOpsqueueConsumerServerResponse, }; -use pyo3::exceptions::{PyBaseException, PyTimeoutError}; +use pyo3::exceptions::{PyBaseException, PyTimeoutError, PyValueError}; use pyo3::{Bound, PyErr, Python, import_exception}; use crate::common; @@ -188,6 +189,12 @@ impl From> for PyErr { } } +impl From> for PyErr { + fn from(value: CError) -> Self { + PyValueError::new_err(value.0.to_string()) + } +} + impl From for CError> { fn from(value: PyErr) -> Self { CError(E::L(FatalPythonException(value))) diff --git a/libs/opsqueue_python/src/producer.rs b/libs/opsqueue_python/src/producer.rs index fcf2b637..ec64acc7 100644 --- a/libs/opsqueue_python/src/producer.rs +++ b/libs/opsqueue_python/src/producer.rs @@ -1,11 +1,11 @@ -use std::{future::IntoFuture, sync::Arc, time::Duration}; - use pyo3::{ create_exception, exceptions::{PyException, PyStopAsyncIteration}, prelude::*, types::PyIterator, }; +use std::time::TryFromFloatSecsError; +use std::{future::IntoFuture, sync::Arc, time::Duration}; use futures::{StreamExt, TryStreamExt, stream::BoxStream}; use opsqueue::{ @@ -431,6 +431,7 @@ impl ProducerClient { PyChunksIter, E![ FatalPythonException, + TryFromFloatSecsError, Elapsed, errors::SubmissionFailed, InternalProducerClientError @@ -440,18 +441,22 @@ impl ProducerClient { self.block_unless_interrupted(async move { let fut = self.stream_completed_submission_chunks(submission_id); match timeout { - Some(duration) => tokio::time::timeout(Duration::from_secs_f64(duration), fut) - .await - .map_err(|err| CError(R(L(err)))) - .and_then(|err| { - err.map_err(|err| match err.0 { - L(err) => CError(L(err)), - R(err) => CError(R(R(err))), + Some(duration) => { + let duration = Duration::try_from_secs_f64(duration) + .map_err(|err| CError(R(L(err))))?; + tokio::time::timeout(duration, fut) + .await + .map_err(|err| CError(R(R(L(err))))) + .and_then(|err| { + err.map_err(|err| match err.0 { + L(err) => CError(L(err)), + R(err) => CError(R(R(R(err)))), + }) }) - }), + } None => fut.await.map_err(|err| match err.0 { L(err) => CError(L(err)), - R(err) => CError(R(R(err))), + R(err) => CError(R(R(R(err)))), }), } }) diff --git a/libs/opsqueue_python/tests/test_roundtrip.py b/libs/opsqueue_python/tests/test_roundtrip.py index 6279ccc0..e0c95687 100644 --- a/libs/opsqueue_python/tests/test_roundtrip.py +++ b/libs/opsqueue_python/tests/test_roundtrip.py @@ -728,7 +728,7 @@ def process_op(x: int) -> int: def test_unpause_and_complete(opsqueue: OpsqueueProcess) -> None: - """Unpausing a paused submission makes it available to consumers again, + """Unpausing a paused submission makes it available to consumers, and it can be completed normally afterwards.""" url = "file:///tmp/opsqueue/test_unpause_and_complete" producer_client = ProducerClient(f"localhost:{opsqueue.port}", url) @@ -751,7 +751,10 @@ def run_consumer() -> None: consumer_client.run_each_op(lambda x: x) with background_process(run_consumer): - producer_client.blocking_stream_completed_submission(submission_id) + producer_client.blocking_stream_completed_submission( + submission_id, + timeout=SUBMISSION_COMPLETED_TIMEOUT, + ) assert isinstance( producer_client.get_submission_status(submission_id), SubmissionStatus.Completed, diff --git a/opsqueue/src/common/chunk.rs b/opsqueue/src/common/chunk.rs index 9fa4e88b..a2239a0f 100644 --- a/opsqueue/src/common/chunk.rs +++ b/opsqueue/src/common/chunk.rs @@ -301,19 +301,33 @@ pub mod db { output_content: Option>, mut conn: impl WriterConnection, ) -> Result<(), E> { - conn.transaction(move |mut tx| { - Box::pin(async move { - complete_chunk_raw(chunk_id, output_content, &mut tx).await?; - crate::common::submission::db::maybe_complete_submission( - chunk_id.submission_id, - &mut tx, - ) - .await + let chunks_moved = conn + .transaction(move |mut tx| { + Box::pin(async move { + let chunks_moved = + complete_chunk_raw(chunk_id, output_content, &mut tx).await?; + if chunks_moved { + 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: \ + completed, failed, or cancelled before. Ignoring.", + chunk_id + ); + } + + Result::>::Ok(chunks_moved) + }) }) - }) - .await?; + .await?; - counter!(crate::prometheus::CHUNKS_COMPLETED_COUNTER).increment(1); + if chunks_moved { + counter!(crate::prometheus::CHUNKS_COMPLETED_COUNTER).increment(1); + } Ok(()) } @@ -327,7 +341,7 @@ pub mod db { chunk_id: ChunkId, output_content: Option>, mut tx: impl WriterConnection, - ) -> sqlx::Result<()> { + ) -> sqlx::Result { let now = chrono::prelude::Utc::now(); let chunk_moved = query!( " @@ -336,8 +350,7 @@ pub mod db { SELECT submission_id, chunk_index, $1, julianday($2) FROM chunks WHERE chunks.submission_id = $3 AND chunks.chunk_index = $4; - DELETE FROM chunks WHERE chunks.submission_id = $5 AND chunks.chunk_index = $6 - RETURNING submission_id, chunk_index; + DELETE FROM chunks WHERE chunks.submission_id = $5 AND chunks.chunk_index = $6; ", output_content, now, @@ -346,9 +359,10 @@ pub mod db { chunk_id.submission_id, chunk_id.chunk_index, ) - .fetch_optional(tx.get_inner()) + .execute(tx.get_inner()) .await? - .is_some(); + .rows_affected() + > 0; // Defense in depth: Above query could be called twice on the same chunk. For instance, // when the server was restarted and the reservations are forgotten, and the same chunk // was reserved again. @@ -373,15 +387,8 @@ pub mod db { ) .fetch_one(tx.get_inner()) .await?; - } else { - tracing::warn!( - "Could not complete chunk {:?} because it was either: \ - completed, failed, or cancelled before. Ignoring.", - chunk_id - ); } - - Ok(()) + Ok(chunk_moved) } /// Increment retries for a chunk, or move it to failed state. @@ -438,7 +445,6 @@ pub mod db { completed, failed, or cancelled before. Ignoring.", chunk_id ); - Ok::<_, sqlx::Error>(false) } }) @@ -806,9 +812,10 @@ pub mod db { #[cfg(feature = "server-logic")] pub mod test { use crate::common::StrategicMetadataMap; - use crate::common::submission::db::insert_submission_raw; + use crate::common::submission::db::{insert_submission, insert_submission_raw}; use crate::common::submission::{Submission, SubmissionStatus}; use crate::db::{Connection as _, WriterPool}; + use std::assert_matches; use super::db::*; use super::*; @@ -932,6 +939,35 @@ pub mod test { } } + #[sqlx::test(migrator = "crate::MIGRATOR")] + pub async fn test_calling_complete_chunk_twice_for_same_chunk_does_not_error( + db: sqlx::SqlitePool, + ) { + let db = WriterPool::new(db); + let mut conn = db.writer_conn().await.unwrap(); + let (submission, chunks) = Submission::from_vec( + vec![Some("foo".into()), Some("bar".into()), Some("baz".into())], + None, + ChunkSize::default(), + ) + .unwrap(); + + let chunk_id = ChunkId { + submission_id: chunks[0].submission_id, + chunk_index: chunks[0].chunk_index, + }; + + insert_submission(submission.clone(), chunks, &mut conn) + .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).await; + assert_matches!(res, Ok(())); + } + #[sqlx::test(migrator = "crate::MIGRATOR")] pub async fn test_fail_chunk(db: sqlx::SqlitePool) { let db = WriterPool::new(db); @@ -957,4 +993,36 @@ pub mod test { assert_eq!(count_chunks_completed(&mut conn).await.unwrap(), 0); assert_eq!(count_chunks_failed(&mut conn).await.unwrap(), 1); } + + #[sqlx::test(migrator = "crate::MIGRATOR")] + pub async fn test_calling_fail_chunk_after_exceeding_retries_does_not_error( + db: sqlx::SqlitePool, + ) { + let db = WriterPool::new(db); + let mut conn = db.writer_conn().await.unwrap(); + let (submission, chunks) = Submission::from_vec( + vec![Some("foo".into()), Some("bar".into()), Some("baz".into())], + None, + ChunkSize::default(), + ) + .unwrap(); + + let chunk_id = ChunkId { + submission_id: chunks[0].submission_id, + chunk_index: chunks[0].chunk_index, + }; + + insert_submission(submission.clone(), chunks, &mut conn) + .await + .expect("insertion failed"); + + let res = retry_or_fail_chunk(chunk_id, "kapot".into(), &mut conn, 2).await; + assert_matches!(res, Ok(false)); + + let res = retry_or_fail_chunk(chunk_id, "kapot".into(), &mut conn, 2).await; + assert_matches!(res, Ok(true)); + + let res = retry_or_fail_chunk(chunk_id, "kapot".into(), &mut conn, 2).await; + assert_matches!(res, Ok(false)); + } } diff --git a/opsqueue/src/common/submission.rs b/opsqueue/src/common/submission.rs index fab26829..bc7f3b32 100644 --- a/opsqueue/src/common/submission.rs +++ b/opsqueue/src/common/submission.rs @@ -526,6 +526,9 @@ pub mod db { Box::pin(async move { unpause_submission_raw(id, &mut tx).await?; super::chunk::db::restore_paused_chunks(id, &mut tx).await?; + // 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(()) }) }) @@ -805,6 +808,24 @@ pub mod db { ) -> Result, DatabaseError> { // NOTE: The order is important here; a concurrent writer could move a submission // from InProgress to Completed/Failed in-between the queries. + // TODO: Rewrite the queries here into a single query using `UNION ALL`. + + let paused_row_opt = submission_status_paused_query(id) + .fetch_optional(conn.get_inner()) + .await?; + if let Some(row) = paused_row_opt { + let paused_submission = SubmissionPaused { + id: row.id, + prefix: row.prefix, + chunks_total: row.chunks_total, + chunks_done: row.chunks_done, + chunk_size: row.chunk_size, + metadata: row.metadata, + strategic_metadata: row.strategic_metadata.0, + otel_trace_carrier: row.otel_trace_carrier, + }; + return Ok(Some(SubmissionStatus::Paused(paused_submission))); + } let submission_row = submission_status_in_progress_query(id) .fetch_optional(conn.get_inner()) @@ -880,23 +901,6 @@ pub mod db { return Ok(Some(SubmissionStatus::Cancelled(cancelled_submission))); } - let paused_row_opt = submission_status_paused_query(id) - .fetch_optional(conn.get_inner()) - .await?; - if let Some(row) = paused_row_opt { - let paused_submission = SubmissionPaused { - id: row.id, - prefix: row.prefix, - chunks_total: row.chunks_total, - chunks_done: row.chunks_done, - chunk_size: row.chunk_size, - metadata: row.metadata, - strategic_metadata: row.strategic_metadata.0, - otel_trace_carrier: row.otel_trace_carrier, - }; - return Ok(Some(SubmissionStatus::Paused(paused_submission))); - } - Ok(None) } @@ -2124,6 +2128,31 @@ pub mod test { assert_eq!(count_chunks_paused(&mut conn).await.unwrap(), 0); } + #[sqlx::test(migrator = "crate::MIGRATOR")] + 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, chunks) = + Submission::from_vec(vec![], None, ChunkSize::default()).unwrap(); + insert_paused_submission(submission.clone(), chunks, &mut conn) + .await + .expect("insertion failed"); + + assert_eq!(count_submissions(&mut conn).await.unwrap(), 0); + assert_eq!(count_submissions_completed(&mut conn).await.unwrap(), 0); + assert_eq!(count_submissions_paused(&mut conn).await.unwrap(), 1); + 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(); + + assert_eq!(count_submissions(&mut conn).await.unwrap(), 0); + assert_eq!(count_submissions_completed(&mut conn).await.unwrap(), 1); + assert_eq!(count_submissions_paused(&mut conn).await.unwrap(), 0); + assert_eq!(count_chunks(&mut conn).await.unwrap(), 0); + assert_eq!(count_chunks_paused(&mut conn).await.unwrap(), 0); + } + #[sqlx::test(migrator = "crate::MIGRATOR")] pub async fn test_cancel_paused_submission(db: sqlx::SqlitePool) { let db = WriterPool::new(db); From 668443e1e13863fecebe7a6c8a9022229b25f591 Mon Sep 17 00:00:00 2001 From: Sem Mulder Date: Tue, 11 Aug 2026 15:46:13 +0200 Subject: [PATCH 08/11] Add delegation support --- Cargo.lock | 73 +- opsqueue/Cargo.toml | 15 +- ...1150000_submissions_external_task.down.sql | 1 + ...811150000_submissions_external_task.up.sql | 6 + opsqueue/opsqueue_example_database_schema.db | Bin 106496 -> 118784 bytes opsqueue/src/common/chunk.rs | 36 +- opsqueue/src/common/submission.rs | 113 +- opsqueue/src/config.rs | 5 + opsqueue/src/consumer/server/mod.rs | 24 +- opsqueue/src/db/mod.rs | 12 + opsqueue/src/delegation/mod.rs | 2 + opsqueue/src/delegation/server.rs | 1196 +++++++++++++++++ opsqueue/src/lib.rs | 1 + opsqueue/src/producer/server.rs | 20 +- opsqueue/src/server.rs | 23 +- workspace-hack/Cargo.toml | 20 +- 16 files changed, 1439 insertions(+), 108 deletions(-) create mode 100644 opsqueue/migrations/20260811150000_submissions_external_task.down.sql create mode 100644 opsqueue/migrations/20260811150000_submissions_external_task.up.sql create mode 100644 opsqueue/src/delegation/mod.rs create mode 100644 opsqueue/src/delegation/server.rs diff --git a/Cargo.lock b/Cargo.lock index 9d5960f2..9a5ac2c2 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -106,6 +106,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" @@ -665,6 +675,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" @@ -1113,6 +1141,12 @@ version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" +[[package]] +name = "hermit-abi" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" + [[package]] name = "hex" version = "0.4.3" @@ -1779,6 +1813,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" @@ -2142,6 +2186,7 @@ dependencies = [ "tokio", "tokio-tungstenite 0.30.0", "tokio-util", + "tower", "tower-http 0.7.0", "tracing", "tracing-opentelemetry", @@ -2149,6 +2194,7 @@ dependencies = [ "url", "uuid", "ux", + "wiremock", "workspace-hack", ] @@ -3233,7 +3279,6 @@ dependencies = [ "log", "memchr", "percent-encoding", - "rustls", "serde", "serde_json", "sha2 0.10.9", @@ -3243,7 +3288,6 @@ dependencies = [ "tokio-stream", "tracing", "url", - "webpki-roots", ] [[package]] @@ -4350,6 +4394,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" @@ -4360,7 +4427,6 @@ checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e" name = "workspace-hack" version = "0.1.0" dependencies = [ - "aws-lc-rs", "base64", "bitflags", "cc", @@ -4384,6 +4450,7 @@ dependencies = [ "opentelemetry-http", "opentelemetry_sdk", "rand 0.9.5", + "regex", "regex-automata", "regex-syntax", "reqwest", diff --git a/opsqueue/Cargo.toml b/opsqueue/Cargo.toml index 310eb645..ae9f2f18 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] @@ -80,6 +80,8 @@ workspace = true [dev-dependencies] insta.workspace = true +wiremock = "0.6.5" +tower = "0.5.3" [features] # Dependencies only in use by the server-logic: @@ -92,8 +94,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/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 3d3cdb8976d9c200f9abdf79f3fce5f15aded4b1..fd7f67a56fe7f64f996c5930fd5946e844664959 100644 GIT binary patch delta 625 zcmZoTz}B#UeS);0ECT}r2*ZHhL>*&B*^LQH_!aoL^ceUT@VoNe<;&y!#oN!T#j}ei zfcqME4%aWPxmT~(nS?!HEb6A9O!-^+uC@xLP%`7g?%+D)UNUbPIEy_#G zQ7B0)&NebIGSfA%&^0ttFfg|=HnB1_NJuU@+Z+*`eau+m{n~k)E9}381wL;x>UsRx zkMsNa_v_Cuz0#1_vEuZRHNETn-oDmmIHy<~&CAXz4Yf}E<&A>PYo)?j75Mnm8Th~P zzvh3yf0h3f|9<{0{44qA@lWRO;IHK`;!oc!XyDAx&Be^h2yrFf<|6$`9~cFvvnnvk zv$069db9kPIPt@FDMiLcMi$n2jJ4b6D=}6wIyJ`1v5VW=Gd6jaBqrsgBKa;J=DT>1 z?{G>;Pi|n6tVdDHrNN~D1PCqhnJFLvCr@`zzYqmKe;^L^@lgo%^9&4i)d4F8*#T7s zQmUX3;u;YGq_8T^Nh~gjFD^+eDJ_mKPR%P(-~p-T(hO8!6Hm3De2_;)AH_lOMhLeQ mqxwrWP#&ho5LJ&MPCYUlCm1+D>1q2iRmP~QMF9$n90UM~(axLz delta 325 zcmZozz}|3xZGyC*Gy?;J6cEFJ?nE79M(K?SOZeq@x#lwPFW`6OyUUly`-``qSBqyC zPXPBd?i{XPTyr-YDpYfAo-38kD#y!T#K8ZJ|26*u{;T|_`1kW~;a|x=kAE_M2Y)Sp z(Plw|0RGKI`jb8|@=a$|V3eOY@!NJOMaD)(7LF6l+}r0XF;+58FV>, mut conn: impl WriterConnection, - ) -> Result<(), E> { - let chunks_moved = conn + ) -> Result> { + let (chunks_moved, completed_submission) = conn .transaction(move |mut tx| { Box::pin(async move { let chunks_moved = complete_chunk_raw(chunk_id, output_content, &mut tx).await?; + + let mut completed_submission = false; if chunks_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: \ @@ -320,7 +323,10 @@ pub mod db { ); } - Result::>::Ok(chunks_moved) + Result::<(bool, bool), E>::Ok(( + chunks_moved, + completed_submission, + )) }) }) .await?; @@ -328,7 +334,7 @@ pub mod db { if chunks_moved { counter!(crate::prometheus::CHUNKS_COMPLETED_COUNTER).increment(1); } - Ok(()) + Ok(completed_submission) } /// This function MUST be called inside a transaction. @@ -657,8 +663,8 @@ pub mod db { submission_id, submission_id, ) - .execute(conn.get_inner()) - .await?; + .execute(conn.get_inner()) + .await?; Ok(()) } @@ -687,8 +693,8 @@ pub mod db { submission_id, submission_id, ) - .execute(conn.get_inner()) - .await?; + .execute(conn.get_inner()) + .await?; counter!(crate::prometheus::CHUNKS_SKIPPED_COUNTER).increment(query_res.rows_affected()); Ok(()) @@ -962,10 +968,10 @@ pub mod test { .expect("insertion failed"); let res = complete_chunk(chunk_id, None, &mut conn).await; - assert_matches!(res, Ok(())); + assert_matches!(res, Ok(false)); let res = complete_chunk(chunk_id, None, &mut conn).await; - assert_matches!(res, Ok(())); + assert_matches!(res, Ok(false)); } #[sqlx::test(migrator = "crate::MIGRATOR")] diff --git a/opsqueue/src/common/submission.rs b/opsqueue/src/common/submission.rs index bc7f3b32..01f9a764 100644 --- a/opsqueue/src/common/submission.rs +++ b/opsqueue/src/common/submission.rs @@ -296,6 +296,11 @@ impl Submission { #[cfg(feature = "server-logic")] pub mod db { + use super::{ + Chunk, ChunkCount, ChunkIndex, DateTime, Duration, E, Metadata, Submission, + SubmissionCancelled, SubmissionCompleted, SubmissionFailed, SubmissionId, SubmissionStatus, + Utc, chunk, + }; use crate::tracing::as_dyn_error; use crate::{ common::{ @@ -312,12 +317,6 @@ pub mod db { use chunk::ChunkSize; use sqlx::{Database, QueryBuilder, Sqlite, query, query_as, query_scalar}; - use super::{ - Chunk, ChunkCount, ChunkIndex, DateTime, Duration, E, Metadata, Submission, - SubmissionCancelled, SubmissionCompleted, SubmissionFailed, SubmissionId, SubmissionStatus, - Utc, chunk, - }; - impl<'q> sqlx::Encode<'q, Sqlite> for SubmissionId { fn encode_by_ref( &self, @@ -379,8 +378,8 @@ pub mod db { submission.otel_trace_carrier, submission.chunk_size.0, ) - .execute(conn.get_inner()) - .await?; + .execute(conn.get_inner()) + .await?; Ok(()) } @@ -503,8 +502,8 @@ pub mod db { submission.otel_trace_carrier, submission.chunk_size.0, ) - .execute(conn.get_inner()) - .await?; + .execute(conn.get_inner()) + .await?; Ok(()) } @@ -1165,6 +1164,9 @@ pub mod db { // 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:?}") } @@ -1177,15 +1179,6 @@ 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)), } } @@ -1200,29 +1193,23 @@ pub mod db { /// # Errors /// /// Returns an error if cancellation or chunk skipping fails. - async fn cancel_submission_notx( + pub(crate) async fn cancel_submission_notx( id: SubmissionId, - mut conn: impl WriterConnection, + mut conn: impl WriterConnection, ) -> Result<(), E> { - cancel_submission_raw(id, &mut conn).await?; - super::chunk::db::skip_remaining_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?; - /// Do not call directly! Must be called inside a transaction. - /// - /// # 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( - 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(()) + Ok(()) + } + Err(E::L(db_err)) => Err(E::L(db_err)), + } } #[tracing::instrument(skip(conn))] @@ -1244,8 +1231,8 @@ pub mod db { id, id, ) - .execute(conn.get_inner()) - .await?; + .execute(conn.get_inner()) + .await?; if res.rows_affected() == 0 { Err(E::R(SubmissionNotFound(id))) } else { @@ -1276,8 +1263,8 @@ pub mod db { id, id, ) - .execute(conn.get_inner()) - .await?; + .execute(conn.get_inner()) + .await?; if res.rows_affected() == 0 { Err(E::R(SubmissionNotFound(id))) } else { @@ -1347,8 +1334,8 @@ pub mod db { id, id, ) - .fetch_one(conn.get_inner()) - .await?; + .fetch_one(conn.get_inner()) + .await?; counter!(crate::prometheus::SUBMISSIONS_FAILED_COUNTER).increment(1); histogram!(crate::prometheus::SUBMISSIONS_DURATION_FAIL_HISTOGRAM).record( crate::prometheus::time_delta_as_f64(Utc::now() - id.timestamp()), @@ -1500,6 +1487,8 @@ pub mod db { tracing::info!("Cleaning up old completed/failed submissions..."); conn.transaction(move |mut tx| { Box::pin(async move { + // TODO(delegation): Prevent deletion if it is still referenced in + // `submissions_external_task`. // Clean up old submissions_metadata query!( "DELETE FROM submissions_metadata @@ -1508,8 +1497,8 @@ pub mod db { );", older_than ) - .execute(tx.get_inner()) - .await?; + .execute(tx.get_inner()) + .await?; query!( "DELETE FROM submissions_metadata WHERE submission_id IN ( @@ -1517,8 +1506,8 @@ pub mod db { );", older_than ) - .execute(tx.get_inner()) - .await?; + .execute(tx.get_inner()) + .await?; query!( "DELETE FROM submissions_metadata WHERE submission_id IN ( @@ -1526,41 +1515,41 @@ pub mod db { );", older_than ) - .execute(tx.get_inner()) - .await?; + .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(); + .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(); + .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(); + .execute(tx.get_inner()) + .await?.rows_affected(); let n_chunks_completed = query!( "DELETE FROM chunks_completed WHERE completed_at < julianday($1);", older_than ) - .execute(tx.get_inner()) - .await?.rows_affected(); + .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(); + .execute(tx.get_inner()) + .await?.rows_affected(); 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)"); @@ -1568,7 +1557,7 @@ pub mod db { Ok(()) }) }) - .await + .await } pub async fn periodically_cleanup_old(db: &WriterPool, max_age: Duration) { 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..e0baab8a 100644 --- a/opsqueue/src/consumer/server/mod.rs +++ b/opsqueue/src/consumer/server/mod.rs @@ -37,10 +37,12 @@ pub async fn serve_for_tests( reservation_expiration: Duration, ) { let notify_on_insert = Arc::new(Notify::new()); + let notify_on_submission_change = Arc::new(Notify::new()); let config = Box::leak(Box::default()); let state = ServerState::new( pool, notify_on_insert, + notify_on_submission_change, cancellation_token.clone(), reservation_expiration, config, @@ -73,13 +75,18 @@ impl ServerState { pub fn new( pool: DBPools, notify_on_insert: Arc, + notify_on_submission_change: Arc, 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, + notify_on_submission_change, + ); Self { pool, completer: Some(completer), @@ -186,6 +193,7 @@ pub struct Completer { dispatcher: Dispatcher, count: usize, max_chunk_retries: u32, + notify_on_submission_change: Arc, } impl Completer { @@ -194,6 +202,7 @@ impl Completer { pool: &db::WriterPool, dispatcher: &Dispatcher, max_chunk_retries: u32, + notify_on_submission_change: Arc, ) -> (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, + notify_on_submission_change, }; (me, tx) } @@ -237,7 +247,7 @@ impl Completer { } => { // Even in the unlikely event that the DB write fails, // we still want to unreserve the chunk - let db_res = + let submission_completed = crate::common::chunk::db::complete_chunk(id, output_content, &mut conn) .await; @@ -259,7 +269,9 @@ impl Completer { let _ = db::perform_explicit_wal_checkpoint(conn).await; } - db_res?; + if submission_completed? { + self.notify_on_submission_change.notify_one(); + } Ok(()) } CompleterMessage::Fail { @@ -293,7 +305,9 @@ impl Completer { histogram!(crate::prometheus::CONSUMER_FAIL_CHUNK_DURATION) .record(start.elapsed()); - failed_permanently?; + if failed_permanently? { + self.notify_on_submission_change.notify_one(); + } Ok(()) } } 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/mod.rs b/opsqueue/src/delegation/mod.rs new file mode 100644 index 00000000..d46f6a49 --- /dev/null +++ b/opsqueue/src/delegation/mod.rs @@ -0,0 +1,2 @@ +#[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..f0868afb --- /dev/null +++ b/opsqueue/src/delegation/server.rs @@ -0,0 +1,1196 @@ +use crate::common::errors::{E, SubmissionNotFound}; +use crate::common::submission::{self, SubmissionId}; +use crate::config::Config; +use crate::db::{Connection, DBPools, WriterConnection}; +use axum::extract::State; +use axum::http::StatusCode; +use axum::routing::post; +use axum::{Json, Router}; +use std::sync::Arc; +use tokio::select; +use tokio::sync::Notify; +use tokio_util::sync::CancellationToken; + +#[cfg(test)] +pub(crate) fn app_for_tests( + pool: DBPools, + cancellation_token: &CancellationToken, + delegation_server_url: url::Url, + notify_on_submission_change: Arc, +) -> Router { + let notify_on_insert = Arc::new(Notify::new()); + let config: &mut Config = Box::leak(Box::default()); + config.delegation_server_url = Some(delegation_server_url); + let router = ServerState::new( + pool, + config, + cancellation_token.clone(), + notify_on_insert, + notify_on_submission_change, + ) + .run_background() + .build_router(); + + Router::new().nest("/job", router) +} + +#[derive(Debug, Clone)] +pub struct ServerState { + pool: DBPools, + cancellation_token: CancellationToken, + /// Notified when new chunks become available for dispatch (e.g. after unpausing a submission). + pub notify_on_insert: Arc, + /// Notified whenever a submission changes status, so the background loop can report + /// it to the external service. + pub notify_on_submission_change: Arc, + delegation_server_url: url::Url, + http_client: reqwest::Client, +} + +impl ServerState { + /// # Panics + /// + /// Panics if `config.delegation_server_url` is not set. + pub fn new( + pool: DBPools, + config: &'static Config, + cancellation_token: CancellationToken, + notify_on_insert: Arc, + notify_on_submission_change: Arc, + ) -> Self { + Self { + pool, + cancellation_token, + notify_on_insert, + notify_on_submission_change, + delegation_server_url: config + .delegation_server_url + .clone() + .expect("delegation_server_url not set"), + http_client: reqwest::Client::new(), + } + } + + #[must_use] + pub fn run_background(self) -> Self { + let state = self.clone(); + let cancellation_token = self.cancellation_token.clone(); + tokio::spawn(async move { + run_in_background( + state.notify_on_submission_change.clone(), + state, + cancellation_token, + ) + .await + .ok(); + }); + self + } + + pub fn build_router(self: ServerState) -> Router<()> { + Router::new() + .route("/delegate", post(job_delegate)) + .route("/kill", post(job_kill)) + .route("/return", post(job_return)) + // .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, +} + +// TODO(delegation): Switch to +// #[tracing::instrument(level = "debug", skip(state))] +// async fn submit( +// State(state): State, +// Json(events): Json>, +// ) -> Result { +// let mut conn = state.pool.writer_conn().await.map_err(|e| { +// tracing::error!("DB error acquiring writer connection: {e:?}"); +// StatusCode::INTERNAL_SERVER_ERROR +// })?; +// // TODO(delegation): Operate within a transaction. +// for event in events { +// match event { +// WorkerDelegationEvent::Delegate(delegations) => { +// for delegation in delegations { +// handle_delegate_event(&state, &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(&state, &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." +// ); +// return Ok(StatusCode::ACCEPTED); +// } +// } +// } +// +// Ok(StatusCode::ACCEPTED) +// } + +#[tracing::instrument(level = "debug", skip(state))] +async fn job_delegate( + State(state): State, + Json(job): Json, +) -> Result { + let mut conn = state.pool.writer_conn().await.map_err(|e| { + tracing::error!("DB error acquiring writer connection: {e:?}"); + StatusCode::INTERNAL_SERVER_ERROR + })?; + handle_delegate_event(&mut conn, &job).await.map_err(|e| { + tracing::error!("DB error handling delegate event: {e:?}"); + StatusCode::INTERNAL_SERVER_ERROR + })?; + + state.notify_on_submission_change.notify_one(); + state.notify_on_insert.notify_waiters(); + + Ok(StatusCode::ACCEPTED) +} + +#[tracing::instrument(level = "debug", skip(state))] +async fn job_kill( + State(state): State, + Json(task_ids): Json>, +) -> Result { + let mut conn = state.pool.writer_conn().await.map_err(|e| { + tracing::error!("DB error acquiring writer connection: {e:?}"); + StatusCode::INTERNAL_SERVER_ERROR + })?; + + conn.transaction(move |mut tx| { + Box::pin(async move { + for task_id in &task_ids { + handle_kill_event(&mut tx, task_id).await?; + } + + Ok::<(), sqlx::Error>(()) + }) + }) + .await + .map_err(|e| { + tracing::error!("DB error handling kill event: {e:?}"); + StatusCode::INTERNAL_SERVER_ERROR + })?; + + state.notify_on_submission_change.notify_one(); + + Ok(StatusCode::ACCEPTED) +} + +#[tracing::instrument(level = "debug", skip(_state))] +async fn job_return( + State(_state): State, + Json(task_ids): Json>, +) -> Result { + tracing::info!("Received 'return' delegation event, which is not yet implemented; ignoring."); + + Ok(StatusCode::ACCEPTED) +} + +#[tracing::instrument(level = "debug", skip(conn))] +async fn handle_delegate_event( + 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 submission::db::unpause_submission(submission_id, &mut *conn).await { + Ok(()) => {} + Err(E::R(SubmissionNotFound(_))) => { + tracing::debug!(%submission_id, "Submission was not in paused state; assuming already active"); + } + 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(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 submission::db::cancel_submission_notx(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(SubmissionNotFound(_))) => { + tracing::warn!(%submission_id, "Submission not found when attempting to cancel; already gone"); + } + } + + Ok(()) +} + +const DELEGATION_BACKGROUND_LOOP_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5); + +async fn run_in_background( + notify_on_submission_change: Arc, + state: ServerState, + cancellation_token: CancellationToken, +) -> Result<(), ()> { + tracing::info!( + "Started delegation background loop. Updates will be sent to {}", + state.delegation_server_url + ); + + let mut triggered_by_timeout: bool = false; + + loop { + match report_submission_status(&state, triggered_by_timeout).await { + Ok(()) => {} + Err(e) => tracing::error!("Error in delegation background loop: {e:?}"), + } + + triggered_by_timeout = select! { + () = cancellation_token.cancelled() => break, + () = notify_on_submission_change.notified() => false, + () = tokio::time::sleep(DELEGATION_BACKGROUND_LOOP_TIMEOUT) => true, + }; + } + + Ok(()) +} + +async fn report_submission_status( + state: &ServerState, + triggered_by_timeout: bool, +) -> anyhow::Result<()> { + let out_of_date_tasks = { + let conn = state.pool.reader_conn().await?; + select_out_of_date_tasks(conn).await? + }; + + if out_of_date_tasks.is_empty() { + return Ok(()); + } + + if triggered_by_timeout { + tracing::warn!( + n_out_of_date_tasks = out_of_date_tasks.len(), + "Delegation background loop triggered by timeout with pending tasks; \ + possible missing notify_on_submission_change call" + ); + } + + for batch in out_of_date_tasks.chunks(2048) { + let mut updates = Vec::new(); + let mut completions = Vec::new(); + + for task in batch { + match task.current_status { + DelegatedJobStatus::Paused => updates.push(DelegatedJobUpdate { + task_id: &task.task_id, + status: DelegatedJobUpdateStatus::Queued, + }), + DelegatedJobStatus::InProgress => updates.push(DelegatedJobUpdate { + task_id: &task.task_id, + status: DelegatedJobUpdateStatus::Running, + }), + DelegatedJobStatus::Completed => completions.push(DelegatedJobCompletion { + task_id: &task.task_id, + completion: DelegatedJobCompletionStatus::Success, + }), + DelegatedJobStatus::Failed => completions.push(DelegatedJobCompletion { + task_id: &task.task_id, + completion: DelegatedJobCompletionStatus::Failure { + failure_reason: FailureReason::Unknown, + }, + }), + DelegatedJobStatus::Cancelled => completions.push(DelegatedJobCompletion { + task_id: &task.task_id, + completion: DelegatedJobCompletionStatus::Failure { + failure_reason: FailureReason::Forced, + }, + }), + } + } + + if !updates.is_empty() { + send_updates(state, &updates).await?; + let conn = state.pool.writer_conn().await?; + update_last_status_sent( + conn, + out_of_date_tasks + .iter() + .filter(|task| { + task.current_status == DelegatedJobStatus::Paused + || task.current_status == DelegatedJobStatus::InProgress + }) + .collect(), + ) + .await?; + } + + if !completions.is_empty() { + send_completions(state, &completions).await?; + let conn = state.pool.writer_conn().await?; + delete_external_tasks( + conn, + out_of_date_tasks + .iter() + .filter(|task| { + task.current_status == DelegatedJobStatus::Completed + || task.current_status == DelegatedJobStatus::Failed + || task.current_status == DelegatedJobStatus::Cancelled + }) + .collect(), + ) + .await?; + } + } + + Ok(()) +} + +async fn insert_external_task( + mut conn: 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 OutOfDateTaskRow { + task_id: String, + current_status: DelegatedJobStatus, +} + +async fn select_out_of_date_tasks( + mut conn: impl Connection, +) -> sqlx::Result> { + sqlx::query_as!( + OutOfDateTaskRow, + r#"WITH out_of_date_tasks AS ( + SELECT + submission_id, + task_id + FROM submissions_external_task as t + WHERE + t.last_status_sent IS NULL + OR (t.last_status_sent = 'paused' AND NOT EXISTS(SELECT * FROM submissions_paused AS s WHERE s.id = t.submission_id)) + OR (t.last_status_sent = 'in_progress' AND NOT EXISTS(SELECT * FROM submissions AS s WHERE s.id = t.submission_id)) + OR (t.last_status_sent = 'completed' AND NOT EXISTS(SELECT * FROM submissions_completed AS s WHERE s.id = t.submission_id)) + OR (t.last_status_sent = 'failed' AND NOT EXISTS(SELECT * FROM submissions_failed AS s WHERE s.id = t.submission_id)) + OR (t.last_status_sent = 'cancelled' AND NOT EXISTS(SELECT * FROM submissions_cancelled AS s WHERE s.id = t.submission_id)) + ) + SELECT + task_id, + coalesce( + (SELECT 'paused' FROM submissions_paused AS s WHERE s.id = t.submission_id), + (SELECT 'in_progress' FROM submissions AS s WHERE s.id = t.submission_id), + (SELECT 'completed' FROM submissions_completed AS s WHERE s.id = t.submission_id), + (SELECT 'failed' FROM submissions_failed AS s WHERE s.id = t.submission_id), + (SELECT 'cancelled' FROM submissions_cancelled AS s WHERE s.id = t.submission_id) + ) AS "current_status!: DelegatedJobStatus" + FROM out_of_date_tasks AS t + "#) + .fetch_all(conn.get_inner()) + .await +} + +async fn update_last_status_sent( + mut conn: impl WriterConnection, + tasks: Vec<&OutOfDateTaskRow>, +) -> sqlx::Result<()> { + let tasks = tasks + .iter() + .map(|t| (t.current_status, t.task_id.clone())) + .collect::>(); + + conn.transaction(move |mut tx| { + Box::pin(async move { + for (current_status, task_id) in tasks { + sqlx::query!( + "UPDATE submissions_external_task SET last_status_sent = $1 WHERE task_id = $2", + current_status, + task_id, + ) + .execute(tx.get_inner()) + .await?; + } + + Ok::<_, sqlx::Error>(()) + }) + }) + .await?; + + Ok(()) +} + +async fn delete_external_tasks( + mut conn: impl WriterConnection, + tasks: Vec<&OutOfDateTaskRow>, +) -> sqlx::Result<()> { + let tasks = tasks.iter().map(|t| t.task_id.clone()).collect::>(); + + conn.transaction(move |mut tx| { + Box::pin(async move { + for task_id in tasks { + sqlx::query!( + "DELETE FROM submissions_external_task WHERE task_id = $1", + task_id, + ) + .execute(tx.get_inner()) + .await?; + } + + Ok::<_, sqlx::Error>(()) + }) + }) + .await?; + + Ok(()) +} + +// TODO(delegation): Replace `send_updates` and `send_completions` with `send_events`, +// after https://github.com/channable/jobmachine/pull/2210 is merged. +// async fn send_events( +// state: &ServerState, +// events: &MasterDelegationEvent<'_>, +// ) -> reqwest::Result<()> { +// state +// .http_client +// .put( +// state +// .delegation_server_url +// .join("/delegation/submit") +// .unwrap(), +// ) +// .json(&events) +// .send() +// .await? +// .error_for_status()?; +// +// Ok(()) +// } + +async fn send_updates( + state: &ServerState, + updates: &[DelegatedJobUpdate<'_>], +) -> reqwest::Result<()> { + state + .http_client + .put( + state + .delegation_server_url + .join("/delegation/update") + .unwrap(), + ) + .json(updates) + .send() + .await? + .error_for_status()?; + + Ok(()) +} + +async fn send_completions( + state: &ServerState, + completions: &[DelegatedJobCompletion<'_>], +) -> reqwest::Result<()> { + state + .http_client + .put( + state + .delegation_server_url + .join("/delegation/complete") + .unwrap(), + ) + .json(completions) + .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::submission::db::{ + cancel_submission, count_submissions, count_submissions_cancelled, + count_submissions_paused, insert_submission_from_chunks, unpause_submission, + }; + use crate::db::{Connection, DBPools}; + use crate::delegation::server::{ + DelegatedJob, DelegatedJobPayload, app_for_tests, insert_external_task, + }; + use axum::body::Body; + use axum::http::Request; + use http::{StatusCode, header}; + use serde_json::json; + use std::sync::{Arc, Mutex}; + use tokio::sync::{Notify, 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")) + } + + #[sqlx::test(migrator = "crate::MIGRATOR")] + pub async fn test_job_delegation( + 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 pool = DBPools::from_test_pools(&reader_pool, &writer_pool); + + let external_server = MockServer::start().await; + + let cancellation_token = CancellationToken::new(); + let app = app_for_tests( + pool.clone(), + &cancellation_token, + external_server.uri().parse().unwrap(), + Arc::new(Notify::new()), + ); + + let submission = { + let mut conn = 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(), + true, + &mut conn, + ) + .await + .unwrap() + }; + + { + let mut conn = 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("PUT")) + .and(path("/delegation/update")) + .and(body_partial_json( + json!([{"task_id": "test", "status": "running"}]), + )) + .respond_with(SignalResponder::new(tx, ResponseTemplate::new(202))) + .expect(1) + .mount(&external_server) + .await; + + let response = app + .clone() + .oneshot( + Request::builder() + .uri("/job/delegate") + .method("POST") + .header(header::CONTENT_TYPE, "application/json") + .body(Body::from( + serde_json::to_string(&DelegatedJob { + task_id: "test".to_string(), + payload: DelegatedJobPayload { + submission_id: submission, + }, + }) + .unwrap(), + )) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!( + response.status(), + StatusCode::ACCEPTED, + "request failed: {response:?}" + ); + + tokio::time::timeout(std::time::Duration::from_secs(2), rx) + .await + .expect("Timed out waiting for HTTP request") + .expect("Sender dropped without signaling"); + + { + let mut conn = 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 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 cancellation_token = CancellationToken::new(); + let app = app_for_tests( + pool.clone(), + &cancellation_token, + external_server.uri().parse().unwrap(), + Arc::new(Notify::new()), + ); + + { + let mut conn = 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(), + true, + &mut conn, + ) + .await + .unwrap(); + + insert_external_task(&mut conn, submission, "test") + .await + .unwrap(); + }; + + let (tx, rx) = oneshot::channel::<()>(); + Mock::given(method("PUT")) + .and(path("/delegation/complete")) + .and(body_partial_json(json!([{"task_id": "test", "completion": {"status": "failure", "failure_reason": "forced"}}]))) + .respond_with(SignalResponder::new(tx, ResponseTemplate::new(202))) + .expect(1) + .mount(&external_server) + .await; + + let response = app + .clone() + .oneshot( + Request::builder() + .uri("/job/kill") + .method("POST") + .header(header::CONTENT_TYPE, "application/json") + .body(Body::from(serde_json::to_string(&["test"]).unwrap())) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!( + response.status(), + StatusCode::ACCEPTED, + "request failed: {response:?}" + ); + + { + let mut conn = pool.reader_conn().await.unwrap(); + assert_eq!(count_submissions_cancelled(&mut conn).await.unwrap(), 1); + assert_eq!(count_submissions(&mut conn).await.unwrap(), 0); + } + + tokio::time::timeout(std::time::Duration::from_secs(2), rx) + .await + .expect("Timed out waiting for HTTP request") + .expect("Sender dropped without signaling"); + + // Wait for background loop to remove external tasks; + tokio::time::sleep(std::time::Duration::from_millis(200)).await; + + { + let mut conn = 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 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 cancellation_token = CancellationToken::new(); + let notify_on_submission_change = Arc::new(Notify::new()); + let _ = app_for_tests( + pool.clone(), + &cancellation_token, + external_server.uri().parse().unwrap(), + notify_on_submission_change.clone(), + ); + + let submission = { + let mut conn = 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(), + true, + &mut conn, + ) + .await + .unwrap(); + + insert_external_task(&mut conn, submission, "test") + .await + .unwrap(); + + submission + }; + + let (tx, rx) = oneshot::channel::<()>(); + Mock::given(method("PUT")) + .and(path("/delegation/update")) + .and(body_partial_json( + json!([{"task_id": "test", "status": "running"}]), + )) + .respond_with(SignalResponder::new(tx, ResponseTemplate::new(202))) + .expect(1) + .mount(&external_server) + .await; + + { + let mut conn = pool.writer_conn().await.unwrap(); + unpause_submission(submission, &mut conn).await.unwrap(); + notify_on_submission_change.notify_one(); + } + + 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_complete_update( + 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 pool = DBPools::from_test_pools(&reader_pool, &writer_pool); + + let external_server = MockServer::start().await; + + let cancellation_token = CancellationToken::new(); + let notify_on_submission_change = Arc::new(Notify::new()); + let _ = app_for_tests( + pool.clone(), + &cancellation_token, + external_server.uri().parse().unwrap(), + notify_on_submission_change.clone(), + ); + + let submission = { + let mut conn = 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(), + false, + &mut conn, + ) + .await + .unwrap(); + + insert_external_task(&mut conn, submission, "test") + .await + .unwrap(); + + submission + }; + + let (tx, rx) = oneshot::channel::<()>(); + Mock::given(method("PUT")) + .and(path("/delegation/complete")) + .and(body_partial_json( + json!([{"task_id": "test", "completion": {"status": "success"}}]), + )) + .respond_with(SignalResponder::new(tx, ResponseTemplate::new(202))) + .expect(1) + .mount(&external_server) + .await; + + { + let mut conn = pool.writer_conn().await.unwrap(); + complete_chunk((submission, ChunkIndex::zero()).into(), None, &mut conn) + .await + .unwrap(); + notify_on_submission_change.notify_one(); + } + + 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_fail_update( + 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 pool = DBPools::from_test_pools(&reader_pool, &writer_pool); + + let external_server = MockServer::start().await; + + let cancellation_token = CancellationToken::new(); + let notify_on_submission_change = Arc::new(Notify::new()); + let _ = app_for_tests( + pool.clone(), + &cancellation_token, + external_server.uri().parse().unwrap(), + notify_on_submission_change.clone(), + ); + + let submission = { + let mut conn = 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(), + false, + &mut conn, + ) + .await + .unwrap(); + + insert_external_task(&mut conn, submission, "test") + .await + .unwrap(); + + submission + }; + + let (tx, rx) = oneshot::channel::<()>(); + Mock::given(method("PUT")) + .and(path("/delegation/complete")) + .and(body_partial_json(json!([{"task_id": "test", "completion": {"status": "failure", "failure_reason": "unknown"}}]))) + .respond_with(SignalResponder::new(tx, ResponseTemplate::new(202))) + .expect(1) + .mount(&external_server) + .await; + + { + let mut conn = pool.writer_conn().await.unwrap(); + retry_or_fail_chunk( + (submission, ChunkIndex::zero()).into(), + "extreme error".to_owned(), + &mut conn, + 0, + ) + .await + .unwrap(); + notify_on_submission_change.notify_one(); + } + + 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_cancel_update( + 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 pool = DBPools::from_test_pools(&reader_pool, &writer_pool); + + let external_server = MockServer::start().await; + + let cancellation_token = CancellationToken::new(); + let notify_on_submission_change = Arc::new(Notify::new()); + let _ = app_for_tests( + pool.clone(), + &cancellation_token, + external_server.uri().parse().unwrap(), + notify_on_submission_change.clone(), + ); + + let submission = { + let mut conn = 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(), + false, + &mut conn, + ) + .await + .unwrap(); + + insert_external_task(&mut conn, submission, "test") + .await + .unwrap(); + + submission + }; + + let (tx, rx) = oneshot::channel::<()>(); + Mock::given(method("PUT")) + .and(path("/delegation/complete")) + .and(body_partial_json(json!([{"task_id": "test", "completion": {"status": "failure", "failure_reason": "forced"}}]))) + .respond_with(SignalResponder::new(tx, ResponseTemplate::new(202))) + .expect(1) + .mount(&external_server) + .await; + + { + let mut conn = pool.writer_conn().await.unwrap(); + cancel_submission(submission, &mut conn).await.unwrap(); + notify_on_submission_change.notify_one(); + } + + tokio::time::timeout(std::time::Duration::from_secs(2), rx) + .await + .expect("Timed out waiting for HTTP request") + .expect("Sender dropped without signaling"); + } +} 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 7e603233..ed85c4d2 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()), + Arc::new(Notify::new()), + max_submissions, + ) + .serve_for_tests(server_addr) + .await; } #[derive(Debug, Clone)] pub struct ServerState { pool: DBPools, notify_on_insert: Arc, + notify_on_submission_change: Arc, max_submissions: MaxSubmissions, } @@ -33,11 +39,13 @@ impl ServerState { pub fn new( pool: DBPools, notify_on_insert: Arc, + notify_on_submission_change: Arc, max_submissions: MaxSubmissions, ) -> Self { ServerState { pool, notify_on_insert, + notify_on_submission_change, max_submissions, } } @@ -131,7 +139,10 @@ async fn cancel_submission( .await .map_err(|e| ServerError(e.into()).into_response())?; match submission::db::cancel_submission(submission_id, &mut conn).await { - Ok(()) => Ok(()), + Ok(()) => { + state.notify_on_submission_change.notify_one(); + 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()) @@ -158,6 +169,7 @@ async fn unpause_submission( Ok(()) => { // Wake up any waiting consumers now that new chunks are available. state.notify_on_insert.notify_waiters(); + state.notify_on_submission_change.notify_one(); Ok(()) } Err(L(db_err)) => Err(ServerError(db_err.into()).into_response()), diff --git a/opsqueue/src/server.rs b/opsqueue/src/server.rs index 5f9ec428..d6f78243 100644 --- a/opsqueue/src/server.rs +++ b/opsqueue/src/server.rs @@ -99,10 +99,12 @@ pub fn build_router( prometheus_config: crate::prometheus::PrometheusConfig, ) -> Router<()> { let notify_on_insert = Arc::new(Notify::new()); + let notify_on_submission_change = Arc::new(Notify::new()); let consumer_routes = crate::consumer::server::ServerState::new( pool.clone(), notify_on_insert.clone(), + notify_on_submission_change.clone(), cancellation_token.clone(), reservation_expiration, config, @@ -110,16 +112,31 @@ pub fn build_router( .run_background() .build_router(); let producer_routes = crate::producer::server::ServerState::new( - pool, - notify_on_insert, + pool.clone(), + notify_on_insert.clone(), + notify_on_submission_change.clone(), config.max_submissions_returned, ) .build_router(); - let routes = Router::new() + let mut routes = Router::new() .nest("/producer", producer_routes) .nest("/consumer", consumer_routes); + if config.delegation_server_url.is_some() { + let delegation_routes = crate::delegation::server::ServerState::new( + pool, + config, + cancellation_token.clone(), + notify_on_insert.clone(), + notify_on_submission_change.clone(), + ) + .run_background() + .build_router(); + + routes = routes.nest("/job", delegation_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 c4bae4e2..81e02add 100644 --- a/workspace-hack/Cargo.toml +++ b/workspace-hack/Cargo.toml @@ -15,7 +15,6 @@ publish = false ### BEGIN HAKARI SECTION [dependencies] -aws-lc-rs = { version = "1", default-features = false, features = ["aws-lc-sys", "prebuilt-nasm"] } base64 = { version = "0.22" } chrono = { version = "0.4", features = ["serde"] } crossbeam-epoch = { version = "0.9" } @@ -27,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 +35,9 @@ opentelemetry = { version = "0.32" } opentelemetry-http = { version = "0.32", features = ["reqwest-blocking"] } opentelemetry_sdk = { version = "0.32", default-features = false, features = ["internal-logs", "logs", "metrics", "rt-tokio", "trace"] } rand = { version = "0.9" } -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"] } @@ -45,12 +46,13 @@ serde_json = { version = "1", features = ["raw_value"] } sha2 = { version = "0.10" } slab = { version = "0.4" } smallvec = { version = "1", default-features = false, features = ["const_new"] } -sqlx-core = { version = "0.9", features = ["_rt-tokio", "_tls-rustls-aws-lc-rs", "any", "chrono", "json", "migrate", "offline"] } +sqlx-core = { version = "0.9", features = ["_rt-tokio", "any", "chrono", "json", "migrate", "offline"] } sqlx-sqlite = { version = "0.9", default-features = false, features = ["any", "bundled", "chrono", "deserialize", "json", "load-extension", "migrate", "offline", "unlock-notify"] } 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"] } @@ -58,7 +60,6 @@ uuid = { version = "1", features = ["fast-rng", "serde", "v4", "v7"] } zerocopy = { version = "0.8", default-features = false, features = ["derive", "simd"] } [build-dependencies] -aws-lc-rs = { version = "1", default-features = false, features = ["aws-lc-sys", "prebuilt-nasm"] } base64 = { version = "0.22" } chrono = { version = "0.4", features = ["serde"] } crossbeam-utils = { version = "0.8" } @@ -71,14 +72,13 @@ futures-util = { version = "0.3", features = ["channel", "io", "sink"] } 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"] } -rustls-pki-types = { version = "1", features = ["std"] } serde = { version = "1", features = ["alloc", "derive", "rc"] } serde_core = { version = "1", features = ["alloc", "rc"] } serde_json = { version = "1", features = ["raw_value"] } sha2 = { version = "0.10" } slab = { version = "0.4" } smallvec = { version = "1", default-features = false, features = ["const_new"] } -sqlx-core = { version = "0.9", features = ["_rt-tokio", "_tls-rustls-aws-lc-rs", "any", "chrono", "json", "migrate", "offline"] } +sqlx-core = { version = "0.9", features = ["_rt-tokio", "any", "chrono", "json", "migrate", "offline"] } sqlx-sqlite = { version = "0.9", default-features = false, features = ["any", "bundled", "chrono", "deserialize", "json", "load-extension", "migrate", "offline", "unlock-notify"] } syn = { version = "3", features = ["full", "visit-mut"] } thiserror = { version = "2" } @@ -90,9 +90,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] From 50114e26469b259811d2ff665367fffe71b0a66a Mon Sep 17 00:00:00 2001 From: Vince van Noort Date: Mon, 31 Aug 2026 16:40:08 +0200 Subject: [PATCH 09/11] Add core interface for delegation Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- opsqueue/src/common/submission.rs | 2 +- opsqueue/src/server/interface.rs | 139 ++++++++++++++++++++++++++++++ 2 files changed, 140 insertions(+), 1 deletion(-) create mode 100644 opsqueue/src/server/interface.rs diff --git a/opsqueue/src/common/submission.rs b/opsqueue/src/common/submission.rs index 01f9a764..9c5ca320 100644 --- a/opsqueue/src/common/submission.rs +++ b/opsqueue/src/common/submission.rs @@ -535,7 +535,7 @@ pub mod db { } #[tracing::instrument(skip(conn))] - pub(super) async fn unpause_submission_raw( + pub(crate) async fn unpause_submission_raw( id: SubmissionId, mut conn: impl WriterConnection, ) -> Result<(), E> { diff --git a/opsqueue/src/server/interface.rs b/opsqueue/src/server/interface.rs new file mode 100644 index 00000000..90e966a8 --- /dev/null +++ b/opsqueue/src/server/interface.rs @@ -0,0 +1,139 @@ +use std::sync::Arc; + +use tokio::sync::{Mutex, mpsc}; + +use crate::{ + common::{ + errors::{DatabaseError, E}, + submission::{self, SubmissionId, SubmissionStatus}, + }, + db::{Connection, DBPools}, +}; + +pub(crate) type SubmissionStatusChangedSender = mpsc::UnboundedSender; + +/// The narrow interface exposed by the Opsqueue core to optional integrations. +#[derive(Clone, Debug)] +pub struct Interface { + pub(crate) pool: DBPools, + submission_status_changed: Arc>>, + status_changed_sender: SubmissionStatusChangedSender, +} + +impl Interface { + pub(crate) fn new( + pool: DBPools, + status_changed_sender: SubmissionStatusChangedSender, + status_changed_receiver: mpsc::UnboundedReceiver, + ) -> Self { + Self { + pool, + submission_status_changed: Arc::new(Mutex::new(status_changed_receiver)), + status_changed_sender, + } + } + + /// Unpauses all given submissions in one transaction. + /// + /// Submissions that are no longer paused are ignored. This makes the operation safe to retry + /// after a delegation request has been delivered more than once. + /// + /// # Errors + /// + /// Returns an error if acquiring a writer connection or updating the database fails. + pub async fn unpause_submissions(&self, ids: Vec) -> Result<(), DatabaseError> { + let mut conn = self.pool.writer_conn().await?; + let changed_ids = conn + .transaction(move |mut tx| { + Box::pin(async move { + let mut changed_ids = Vec::new(); + for id in ids { + let was_paused = + match submission::db::unpause_submission_raw(id, &mut tx).await { + Ok(()) => { + changed_ids.push(id); + true + } + Err(E::R(_)) => false, + Err(E::L(error)) => return Err(error), + }; + if was_paused { + crate::common::chunk::db::restore_paused_chunks(id, &mut tx).await?; + match submission::db::maybe_complete_submission(id, &mut tx).await { + Ok(_) | Err(E::R(_)) => {} + Err(E::L(error)) => return Err(error), + } + } + } + Ok(changed_ids) + }) + }) + .await?; + + self.notify_status_changed(changed_ids); + Ok(()) + } + + /// Cancels all given submissions in one transaction. + /// + /// Submissions that are already terminal or missing are ignored. This makes the operation safe + /// to retry after a delegation request has been delivered more than once. + /// + /// # Errors + /// + /// Returns an error if acquiring a writer connection or updating the database fails. + pub async fn cancel_submissions(&self, ids: Vec) -> Result<(), DatabaseError> { + let mut conn = self.pool.writer_conn().await?; + let changed_ids = conn + .transaction(move |mut tx| { + Box::pin(async move { + let mut changed_ids = Vec::new(); + for id in ids { + match submission::db::cancel_submission_notx(id, &mut tx).await { + Ok(()) => changed_ids.push(id), + Err(E::R(_)) => {} + Err(E::L(error)) => return Err(error), + } + } + Ok(changed_ids) + }) + }) + .await?; + + self.notify_status_changed(changed_ids); + Ok(()) + } + + /// Gets the current status for each given submission ID. + /// + /// The returned vector has the same order as `ids`; a `None` entry means that the submission + /// no longer exists. + /// + /// # Errors + /// + /// Returns an error if acquiring a reader connection or querying the database fails. + pub async fn get_submission_statuses( + &self, + ids: Vec, + ) -> Result>, DatabaseError> { + let mut conn = self.pool.reader_conn().await?; + let mut statuses = Vec::with_capacity(ids.len()); + for id in ids { + statuses.push(submission::db::submission_status(id, &mut conn).await?); + } + Ok(statuses) + } + + /// Waits until any submission status changes. + /// + /// Returns `None` if the underlying channel is closed, which happens during shutdown. + pub async fn wait_for_submission_status_change(&self) -> Option { + self.submission_status_changed.lock().await.recv().await + } + + fn notify_status_changed(&self, ids: Vec) { + for id in ids { + let _ = self.status_changed_sender.send(id); + } + } +} From 9216a8c95f99de46a468d66fae1c5473a28c8b4a Mon Sep 17 00:00:00 2001 From: Vince van Noort Date: Mon, 31 Aug 2026 16:40:13 +0200 Subject: [PATCH 10/11] Publish submission status changes to integrations Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- opsqueue/src/consumer/server/mod.rs | 38 ++++++++++++++++++++--------- opsqueue/src/producer/server.rs | 34 ++++++++++++++++---------- opsqueue/src/server.rs | 31 ++++++++++++++--------- 3 files changed, 68 insertions(+), 35 deletions(-) diff --git a/opsqueue/src/consumer/server/mod.rs b/opsqueue/src/consumer/server/mod.rs index e0baab8a..1db3b6df 100644 --- a/opsqueue/src/consumer/server/mod.rs +++ b/opsqueue/src/consumer/server/mod.rs @@ -18,6 +18,7 @@ use crate::{ common::chunk::ChunkId, config::Config, db::{self, DBPools}, + server::interface::SubmissionStatusChangedSender, }; use super::dispatcher::Dispatcher; @@ -37,12 +38,10 @@ pub async fn serve_for_tests( reservation_expiration: Duration, ) { let notify_on_insert = Arc::new(Notify::new()); - let notify_on_submission_change = Arc::new(Notify::new()); let config = Box::leak(Box::default()); let state = ServerState::new( pool, notify_on_insert, - notify_on_submission_change, cancellation_token.clone(), reservation_expiration, config, @@ -75,17 +74,34 @@ impl ServerState { pub fn new( pool: DBPools, notify_on_insert: Arc, - notify_on_submission_change: Arc, cancellation_token: CancellationToken, reservation_expiration: Duration, config: &'static Config, + ) -> Self { + Self::new_with_status_sender( + pool, + notify_on_insert, + cancellation_token, + reservation_expiration, + config, + None, + ) + } + + pub fn new_with_status_sender( + pool: DBPools, + notify_on_insert: Arc, + cancellation_token: CancellationToken, + reservation_expiration: Duration, + config: &'static Config, + status_changed_sender: Option, ) -> Self { let dispatcher = Dispatcher::new(reservation_expiration); let (completer, completer_tx) = Completer::new( pool.writer_pool(), &dispatcher, config.max_chunk_retries, - notify_on_submission_change, + status_changed_sender, ); Self { pool, @@ -193,7 +209,7 @@ pub struct Completer { dispatcher: Dispatcher, count: usize, max_chunk_retries: u32, - notify_on_submission_change: Arc, + status_changed_sender: Option, } impl Completer { @@ -202,7 +218,7 @@ impl Completer { pool: &db::WriterPool, dispatcher: &Dispatcher, max_chunk_retries: u32, - notify_on_submission_change: Arc, + status_changed_sender: Option, ) -> (Self, tokio::sync::mpsc::Sender) { let (tx, rx) = tokio::sync::mpsc::channel(1024); let pool = pool.clone(); @@ -212,7 +228,7 @@ impl Completer { dispatcher: dispatcher.clone(), count: 0, max_chunk_retries, - notify_on_submission_change, + status_changed_sender, }; (me, tx) } @@ -269,8 +285,8 @@ impl Completer { let _ = db::perform_explicit_wal_checkpoint(conn).await; } - if submission_completed? { - self.notify_on_submission_change.notify_one(); + if submission_completed? && let Some(sender) = &self.status_changed_sender { + let _ = sender.send(id.submission_id); } Ok(()) } @@ -305,8 +321,8 @@ impl Completer { histogram!(crate::prometheus::CONSUMER_FAIL_CHUNK_DURATION) .record(start.elapsed()); - if failed_permanently? { - self.notify_on_submission_change.notify_one(); + if failed_permanently? && let Some(sender) = &self.status_changed_sender { + let _ = sender.send(id.submission_id); } Ok(()) } diff --git a/opsqueue/src/producer/server.rs b/opsqueue/src/producer/server.rs index ed85c4d2..ed5994dc 100644 --- a/opsqueue/src/producer/server.rs +++ b/opsqueue/src/producer/server.rs @@ -4,6 +4,7 @@ use crate::common::errors::E::{L, R}; use crate::common::submission::{self, SubmissionId}; use crate::common::{MaxSubmissions, StrategicMetadataMap}; use crate::db::{self, DBPools}; +use crate::server::interface::SubmissionStatusChangedSender; use crate::tracing::anyhow_as_dyn_error; use axum::extract; use axum::extract::{Path, State}; @@ -17,36 +18,39 @@ 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()), - Arc::new(Notify::new()), - max_submissions, - ) - .serve_for_tests(server_addr) - .await; + ServerState::new(database_pool, Arc::new(Notify::new()), max_submissions) + .serve_for_tests(server_addr) + .await; } #[derive(Debug, Clone)] pub struct ServerState { pool: DBPools, notify_on_insert: Arc, - notify_on_submission_change: Arc, max_submissions: MaxSubmissions, + status_changed_sender: Option, } impl ServerState { pub fn new( pool: DBPools, notify_on_insert: Arc, - notify_on_submission_change: Arc, max_submissions: MaxSubmissions, + ) -> Self { + Self::new_with_status_sender(pool, notify_on_insert, max_submissions, None) + } + + pub fn new_with_status_sender( + pool: DBPools, + notify_on_insert: Arc, + max_submissions: MaxSubmissions, + status_changed_sender: Option, ) -> Self { ServerState { pool, notify_on_insert, - notify_on_submission_change, max_submissions, + status_changed_sender, } } @@ -140,7 +144,9 @@ async fn cancel_submission( .map_err(|e| ServerError(e.into()).into_response())?; match submission::db::cancel_submission(submission_id, &mut conn).await { Ok(()) => { - state.notify_on_submission_change.notify_one(); + if let Some(sender) = &state.status_changed_sender { + let _ = sender.send(submission_id); + } Ok(()) } Err(L(db_err)) => Err(ServerError(db_err.into()).into_response()), @@ -169,7 +175,9 @@ async fn unpause_submission( Ok(()) => { // Wake up any waiting consumers now that new chunks are available. state.notify_on_insert.notify_waiters(); - state.notify_on_submission_change.notify_one(); + if let Some(sender) = &state.status_changed_sender { + let _ = sender.send(submission_id); + } Ok(()) } Err(L(db_err)) => Err(ServerError(db_err.into()).into_response()), diff --git a/opsqueue/src/server.rs b/opsqueue/src/server.rs index d6f78243..d7d127c9 100644 --- a/opsqueue/src/server.rs +++ b/opsqueue/src/server.rs @@ -16,6 +16,11 @@ use tokio::select; use tokio::sync::Notify; use tokio_util::sync::CancellationToken; +#[cfg(feature = "server-logic")] +pub mod interface; +#[cfg(feature = "server-logic")] +pub use interface::Interface; + fn retry_policy() -> impl BackoffBuilder { FibonacciBuilder::default() .with_jitter() @@ -44,7 +49,7 @@ pub async fn serve_producer_and_consumer( (|| async { let router = build_router( config, - pool.clone(), + pool, reservation_expiration, cancellation_token, app_healthy_flag.clone(), @@ -92,30 +97,35 @@ pub async fn serve_producer_and_consumer( #[cfg(feature = "server-logic")] pub fn build_router( config: &'static crate::config::Config, - pool: DBPools, + pool: &DBPools, reservation_expiration: Duration, cancellation_token: &CancellationToken, app_healthy_flag: Arc, prometheus_config: crate::prometheus::PrometheusConfig, ) -> Router<()> { let notify_on_insert = Arc::new(Notify::new()); - let notify_on_submission_change = Arc::new(Notify::new()); + let (status_changed_sender, status_changed_receiver) = tokio::sync::mpsc::unbounded_channel(); + let interface = interface::Interface::new( + pool.clone(), + status_changed_sender.clone(), + status_changed_receiver, + ); - let consumer_routes = crate::consumer::server::ServerState::new( + let consumer_routes = crate::consumer::server::ServerState::new_with_status_sender( pool.clone(), notify_on_insert.clone(), - notify_on_submission_change.clone(), cancellation_token.clone(), reservation_expiration, config, + Some(status_changed_sender.clone()), ) .run_background() .build_router(); - let producer_routes = crate::producer::server::ServerState::new( + let producer_routes = crate::producer::server::ServerState::new_with_status_sender( pool.clone(), notify_on_insert.clone(), - notify_on_submission_change.clone(), config.max_submissions_returned, + Some(status_changed_sender), ) .build_router(); @@ -123,13 +133,12 @@ pub fn build_router( .nest("/producer", producer_routes) .nest("/consumer", consumer_routes); - if config.delegation_server_url.is_some() { + if let Some(delegation_server_url) = config.delegation_server_url.clone() { let delegation_routes = crate::delegation::server::ServerState::new( - pool, - config, + delegation_server_url, cancellation_token.clone(), + interface, notify_on_insert.clone(), - notify_on_submission_change.clone(), ) .run_background() .build_router(); From 179617fa3a5ab5ba84d595a37d475df0334e38cc Mon Sep 17 00:00:00 2001 From: Vince van Noort Date: Mon, 31 Aug 2026 16:40:17 +0200 Subject: [PATCH 11/11] Route delegation through the core interface Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- opsqueue/src/delegation/server.rs | 327 +++++++++++++----------------- 1 file changed, 144 insertions(+), 183 deletions(-) diff --git a/opsqueue/src/delegation/server.rs b/opsqueue/src/delegation/server.rs index f0868afb..40d851d3 100644 --- a/opsqueue/src/delegation/server.rs +++ b/opsqueue/src/delegation/server.rs @@ -1,7 +1,8 @@ -use crate::common::errors::{E, SubmissionNotFound}; -use crate::common::submission::{self, SubmissionId}; -use crate::config::Config; -use crate::db::{Connection, DBPools, WriterConnection}; +use crate::common::submission::{SubmissionId, SubmissionStatus}; +#[cfg(test)] +use crate::db::DBPools; +use crate::db::{Connection, WriterConnection}; +use crate::server::interface::Interface; use axum::extract::State; use axum::http::StatusCode; use axum::routing::post; @@ -13,60 +14,50 @@ use tokio_util::sync::CancellationToken; #[cfg(test)] pub(crate) fn app_for_tests( - pool: DBPools, + pool: &DBPools, cancellation_token: &CancellationToken, delegation_server_url: url::Url, - notify_on_submission_change: Arc, -) -> Router { +) -> (Router, tokio::sync::mpsc::UnboundedSender) { let notify_on_insert = Arc::new(Notify::new()); - let config: &mut Config = Box::leak(Box::default()); - config.delegation_server_url = Some(delegation_server_url); + let (status_changed_sender, status_changed_receiver) = tokio::sync::mpsc::unbounded_channel(); let router = ServerState::new( - pool, - config, + delegation_server_url, cancellation_token.clone(), + Interface::new( + pool.clone(), + status_changed_sender.clone(), + status_changed_receiver, + ), notify_on_insert, - notify_on_submission_change, ) .run_background() .build_router(); - Router::new().nest("/job", router) + (Router::new().nest("/job", router), status_changed_sender) } #[derive(Debug, Clone)] pub struct ServerState { - pool: DBPools, + interface: Interface, cancellation_token: CancellationToken, /// Notified when new chunks become available for dispatch (e.g. after unpausing a submission). pub notify_on_insert: Arc, - /// Notified whenever a submission changes status, so the background loop can report - /// it to the external service. - pub notify_on_submission_change: Arc, delegation_server_url: url::Url, http_client: reqwest::Client, } impl ServerState { - /// # Panics - /// - /// Panics if `config.delegation_server_url` is not set. pub fn new( - pool: DBPools, - config: &'static Config, + delegation_server_url: url::Url, cancellation_token: CancellationToken, + interface: Interface, notify_on_insert: Arc, - notify_on_submission_change: Arc, ) -> Self { Self { - pool, + interface, cancellation_token, notify_on_insert, - notify_on_submission_change, - delegation_server_url: config - .delegation_server_url - .clone() - .expect("delegation_server_url not set"), + delegation_server_url, http_client: reqwest::Client::new(), } } @@ -76,13 +67,7 @@ impl ServerState { let state = self.clone(); let cancellation_token = self.cancellation_token.clone(); tokio::spawn(async move { - run_in_background( - state.notify_on_submission_change.clone(), - state, - cancellation_token, - ) - .await - .ok(); + run_in_background(state, cancellation_token).await.ok(); }); self } @@ -224,16 +209,27 @@ async fn job_delegate( State(state): State, Json(job): Json, ) -> Result { - let mut conn = state.pool.writer_conn().await.map_err(|e| { + let mut conn = state.interface.pool.writer_conn().await.map_err(|e| { tracing::error!("DB error acquiring writer connection: {e:?}"); StatusCode::INTERNAL_SERVER_ERROR })?; - handle_delegate_event(&mut conn, &job).await.map_err(|e| { - tracing::error!("DB error handling delegate event: {e:?}"); - StatusCode::INTERNAL_SERVER_ERROR - })?; + insert_external_task(&mut conn, job.payload.submission_id, &job.task_id) + .await + .map_err(|e| { + tracing::error!("DB error handling delegate event: {e:?}"); + StatusCode::INTERNAL_SERVER_ERROR + })?; + drop(conn); + + state + .interface + .unpause_submissions(vec![job.payload.submission_id]) + .await + .map_err(|e| { + tracing::error!("DB error handling delegate event: {e:?}"); + StatusCode::INTERNAL_SERVER_ERROR + })?; - state.notify_on_submission_change.notify_one(); state.notify_on_insert.notify_waiters(); Ok(StatusCode::ACCEPTED) @@ -244,27 +240,42 @@ async fn job_kill( State(state): State, Json(task_ids): Json>, ) -> Result { - let mut conn = state.pool.writer_conn().await.map_err(|e| { + let mut conn = state.interface.pool.reader_conn().await.map_err(|e| { tracing::error!("DB error acquiring writer connection: {e:?}"); StatusCode::INTERNAL_SERVER_ERROR })?; - conn.transaction(move |mut tx| { - Box::pin(async move { - for task_id in &task_ids { - handle_kill_event(&mut tx, task_id).await?; - } - - Ok::<(), sqlx::Error>(()) - }) - }) - .await - .map_err(|e| { - tracing::error!("DB error handling kill event: {e:?}"); - StatusCode::INTERNAL_SERVER_ERROR - })?; + let mut submission_ids = Vec::with_capacity(task_ids.len()); + for task_id in &task_ids { + 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 + .map_err(|e| { + tracing::error!(%task_id, "DB error looking up task for kill event: {e:?}"); + StatusCode::INTERNAL_SERVER_ERROR + })?; + + if let Some(submission_id) = submission_id { + submission_ids.push(submission_id); + } else { + tracing::warn!(%task_id, "Kill event for unknown task_id; ignoring"); + } + } + drop(conn); - state.notify_on_submission_change.notify_one(); + state + .interface + .cancel_submissions(submission_ids) + .await + .map_err(|e| { + tracing::error!("DB error handling kill event: {e:?}"); + StatusCode::INTERNAL_SERVER_ERROR + })?; Ok(StatusCode::ACCEPTED) } @@ -279,68 +290,9 @@ async fn job_return( Ok(StatusCode::ACCEPTED) } -#[tracing::instrument(level = "debug", skip(conn))] -async fn handle_delegate_event( - 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 submission::db::unpause_submission(submission_id, &mut *conn).await { - Ok(()) => {} - Err(E::R(SubmissionNotFound(_))) => { - tracing::debug!(%submission_id, "Submission was not in paused state; assuming already active"); - } - 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(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 submission::db::cancel_submission_notx(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(SubmissionNotFound(_))) => { - tracing::warn!(%submission_id, "Submission not found when attempting to cancel; already gone"); - } - } - - Ok(()) -} - const DELEGATION_BACKGROUND_LOOP_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5); async fn run_in_background( - notify_on_submission_change: Arc, state: ServerState, cancellation_token: CancellationToken, ) -> Result<(), ()> { @@ -359,7 +311,7 @@ async fn run_in_background( triggered_by_timeout = select! { () = cancellation_token.cancelled() => break, - () = notify_on_submission_change.notified() => false, + Some(_) = state.interface.wait_for_submission_status_change() => false, () = tokio::time::sleep(DELEGATION_BACKGROUND_LOOP_TIMEOUT) => true, }; } @@ -371,10 +323,7 @@ async fn report_submission_status( state: &ServerState, triggered_by_timeout: bool, ) -> anyhow::Result<()> { - let out_of_date_tasks = { - let conn = state.pool.reader_conn().await?; - select_out_of_date_tasks(conn).await? - }; + let out_of_date_tasks = select_out_of_date_tasks(&state.interface).await?; if out_of_date_tasks.is_empty() { return Ok(()); @@ -384,7 +333,7 @@ async fn report_submission_status( tracing::warn!( n_out_of_date_tasks = out_of_date_tasks.len(), "Delegation background loop triggered by timeout with pending tasks; \ - possible missing notify_on_submission_change call" + possible missing submission status notification" ); } @@ -423,7 +372,7 @@ async fn report_submission_status( if !updates.is_empty() { send_updates(state, &updates).await?; - let conn = state.pool.writer_conn().await?; + let conn = state.interface.pool.writer_conn().await?; update_last_status_sent( conn, out_of_date_tasks @@ -439,7 +388,7 @@ async fn report_submission_status( if !completions.is_empty() { send_completions(state, &completions).await?; - let conn = state.pool.writer_conn().await?; + let conn = state.interface.pool.writer_conn().await?; delete_external_tasks( conn, out_of_date_tasks @@ -487,37 +436,59 @@ struct OutOfDateTaskRow { current_status: DelegatedJobStatus, } -async fn select_out_of_date_tasks( - mut conn: impl Connection, -) -> sqlx::Result> { - sqlx::query_as!( - OutOfDateTaskRow, - r#"WITH out_of_date_tasks AS ( - SELECT - submission_id, - task_id - FROM submissions_external_task as t - WHERE - t.last_status_sent IS NULL - OR (t.last_status_sent = 'paused' AND NOT EXISTS(SELECT * FROM submissions_paused AS s WHERE s.id = t.submission_id)) - OR (t.last_status_sent = 'in_progress' AND NOT EXISTS(SELECT * FROM submissions AS s WHERE s.id = t.submission_id)) - OR (t.last_status_sent = 'completed' AND NOT EXISTS(SELECT * FROM submissions_completed AS s WHERE s.id = t.submission_id)) - OR (t.last_status_sent = 'failed' AND NOT EXISTS(SELECT * FROM submissions_failed AS s WHERE s.id = t.submission_id)) - OR (t.last_status_sent = 'cancelled' AND NOT EXISTS(SELECT * FROM submissions_cancelled AS s WHERE s.id = t.submission_id)) - ) - SELECT - task_id, - coalesce( - (SELECT 'paused' FROM submissions_paused AS s WHERE s.id = t.submission_id), - (SELECT 'in_progress' FROM submissions AS s WHERE s.id = t.submission_id), - (SELECT 'completed' FROM submissions_completed AS s WHERE s.id = t.submission_id), - (SELECT 'failed' FROM submissions_failed AS s WHERE s.id = t.submission_id), - (SELECT 'cancelled' FROM submissions_cancelled AS s WHERE s.id = t.submission_id) - ) AS "current_status!: DelegatedJobStatus" - FROM out_of_date_tasks AS t - "#) +async fn select_out_of_date_tasks(interface: &Interface) -> anyhow::Result> { + let external_tasks = { + let mut conn = interface.pool.reader_conn().await?; + sqlx::query_as!( + ExternalTaskRow, + r#"SELECT + task_id, + submission_id AS "submission_id: SubmissionId", + last_status_sent AS "last_status_sent: DelegatedJobStatus" + FROM submissions_external_task"# + ) .fetch_all(conn.get_inner()) - .await + .await? + }; + + let submission_ids = external_tasks + .iter() + .map(|task| task.submission_id) + .collect::>(); + let statuses = interface.get_submission_statuses(submission_ids).await?; + + let mut out_of_date_tasks = Vec::new(); + for (task, status) in external_tasks.into_iter().zip(statuses) { + let status = status.ok_or_else(|| { + anyhow::anyhow!( + "Submission {} for external task {} no longer exists", + task.submission_id, + task.task_id + ) + })?; + let current_status = match status { + SubmissionStatus::Paused(_) => DelegatedJobStatus::Paused, + SubmissionStatus::InProgress(_) => DelegatedJobStatus::InProgress, + SubmissionStatus::Completed(_) => DelegatedJobStatus::Completed, + SubmissionStatus::Failed(_, _) => DelegatedJobStatus::Failed, + SubmissionStatus::Cancelled(_) => DelegatedJobStatus::Cancelled, + }; + + if task.last_status_sent != Some(current_status) { + out_of_date_tasks.push(OutOfDateTaskRow { + task_id: task.task_id, + current_status, + }); + } + } + Ok(out_of_date_tasks) +} + +#[derive(Debug)] +struct ExternalTaskRow { + task_id: String, + submission_id: SubmissionId, + last_status_sent: Option, } async fn update_last_status_sent( @@ -654,8 +625,8 @@ pub mod test { use axum::http::Request; use http::{StatusCode, header}; use serde_json::json; - use std::sync::{Arc, Mutex}; - use tokio::sync::{Notify, oneshot}; + use std::sync::Mutex; + use tokio::sync::oneshot; use tokio_util::sync::CancellationToken; use tower::ServiceExt; use wiremock::matchers::{body_partial_json, method, path}; @@ -714,11 +685,10 @@ pub mod test { let external_server = MockServer::start().await; let cancellation_token = CancellationToken::new(); - let app = app_for_tests( - pool.clone(), + let (app, _status_changed_sender) = app_for_tests( + &pool, &cancellation_token, external_server.uri().parse().unwrap(), - Arc::new(Notify::new()), ); let submission = { @@ -814,11 +784,10 @@ pub mod test { let external_server = MockServer::start().await; let cancellation_token = CancellationToken::new(); - let app = app_for_tests( - pool.clone(), + let (app, _status_changed_sender) = app_for_tests( + &pool, &cancellation_token, external_server.uri().parse().unwrap(), - Arc::new(Notify::new()), ); { @@ -910,12 +879,10 @@ pub mod test { let external_server = MockServer::start().await; let cancellation_token = CancellationToken::new(); - let notify_on_submission_change = Arc::new(Notify::new()); - let _ = app_for_tests( - pool.clone(), + let (_app, status_changed_sender) = app_for_tests( + &pool, &cancellation_token, external_server.uri().parse().unwrap(), - notify_on_submission_change.clone(), ); let submission = { @@ -955,7 +922,7 @@ pub mod test { { let mut conn = pool.writer_conn().await.unwrap(); unpause_submission(submission, &mut conn).await.unwrap(); - notify_on_submission_change.notify_one(); + status_changed_sender.send(submission).unwrap(); } tokio::time::timeout(std::time::Duration::from_secs(2), rx) @@ -985,12 +952,10 @@ pub mod test { let external_server = MockServer::start().await; let cancellation_token = CancellationToken::new(); - let notify_on_submission_change = Arc::new(Notify::new()); - let _ = app_for_tests( - pool.clone(), + let (_app, status_changed_sender) = app_for_tests( + &pool, &cancellation_token, external_server.uri().parse().unwrap(), - notify_on_submission_change.clone(), ); let submission = { @@ -1032,7 +997,7 @@ pub mod test { complete_chunk((submission, ChunkIndex::zero()).into(), None, &mut conn) .await .unwrap(); - notify_on_submission_change.notify_one(); + status_changed_sender.send(submission).unwrap(); } tokio::time::timeout(std::time::Duration::from_secs(2), rx) @@ -1062,12 +1027,10 @@ pub mod test { let external_server = MockServer::start().await; let cancellation_token = CancellationToken::new(); - let notify_on_submission_change = Arc::new(Notify::new()); - let _ = app_for_tests( - pool.clone(), + let (_app, status_changed_sender) = app_for_tests( + &pool, &cancellation_token, external_server.uri().parse().unwrap(), - notify_on_submission_change.clone(), ); let submission = { @@ -1112,7 +1075,7 @@ pub mod test { ) .await .unwrap(); - notify_on_submission_change.notify_one(); + status_changed_sender.send(submission).unwrap(); } tokio::time::timeout(std::time::Duration::from_secs(2), rx) @@ -1142,12 +1105,10 @@ pub mod test { let external_server = MockServer::start().await; let cancellation_token = CancellationToken::new(); - let notify_on_submission_change = Arc::new(Notify::new()); - let _ = app_for_tests( - pool.clone(), + let (_app, status_changed_sender) = app_for_tests( + &pool, &cancellation_token, external_server.uri().parse().unwrap(), - notify_on_submission_change.clone(), ); let submission = { @@ -1185,7 +1146,7 @@ pub mod test { { let mut conn = pool.writer_conn().await.unwrap(); cancel_submission(submission, &mut conn).await.unwrap(); - notify_on_submission_change.notify_one(); + status_changed_sender.send(submission).unwrap(); } tokio::time::timeout(std::time::Duration::from_secs(2), rx)