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/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/python/opsqueue/producer.py b/libs/opsqueue_python/python/opsqueue/producer.py index 82a877e3..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", ] @@ -96,6 +98,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 +119,7 @@ def run_submission( metadata=metadata, strategic_metadata=strategic_metadata, chunk_size=chunk_size, + timeout=timeout, ) return _unchunk_iterator(results_iter, serialization_format) @@ -146,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, @@ -162,6 +167,7 @@ def insert_submission( metadata=metadata, strategic_metadata=strategic_metadata, chunk_size=chunk_size, + paused=paused, ) def blocking_stream_completed_submission( @@ -169,6 +175,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 +188,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 +218,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 +237,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, @@ -259,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, @@ -275,10 +284,13 @@ 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( - 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 +301,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 @@ -326,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. @@ -337,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/errors.rs b/libs/opsqueue_python/src/errors.rs index 45f0f7da..043ac11f 100644 --- a/libs/opsqueue_python/src/errors.rs +++ b/libs/opsqueue_python/src/errors.rs @@ -1,17 +1,16 @@ /// 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::chunk::ChunkId; use opsqueue::common::errors::{ - ChunkNotFound, E, IncorrectUsage, SubmissionNotCancellable, SubmissionNotFound, - TooManyMatchingSubmissions, UnexpectedOpsqueueConsumerServerResponse, + E, IncorrectUsage, SubmissionNotCancellable, SubmissionNotFound, TooManyMatchingSubmissions, + UnexpectedOpsqueueConsumerServerResponse, }; -use pyo3::exceptions::PyBaseException; +use pyo3::exceptions::{PyBaseException, PyTimeoutError, PyValueError}; use pyo3::{Bound, PyErr, Python, import_exception}; use crate::common; -use crate::common::{ChunkIndex, SubmissionId}; // Expected errors: import_exception!(opsqueue.exceptions, SubmissionFailedError); @@ -19,7 +18,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 +171,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()) @@ -201,6 +183,18 @@ impl From> for PyErr { } } +impl From> for PyErr { + fn from(_value: CError) -> Self { + PyTimeoutError::new_err("timeout was reached") + } +} + +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/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 e3167a9b..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::{ @@ -18,6 +18,7 @@ use opsqueue::{ producer::client::{Client as ActualClient, InternalProducerClientError}, tracing::CarrierMap, }; +use tokio::time::error::Elapsed; use ux::u63; use crate::{ @@ -158,6 +159,36 @@ impl ProducerClient { }) } + /// Unpause a paused submission, making it available to consumers. + /// + /// 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, @@ -246,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<'_>, @@ -254,6 +285,7 @@ impl ProducerClient { metadata: Option, chunk_size: Option, otel_trace_carrier: CarrierMap, + paused: bool, ) -> CPyResult> { let strategic_metadata = std::collections::HashMap::default(); @@ -265,6 +297,7 @@ impl ProducerClient { }, metadata, strategic_metadata, + paused, }; self.block_unless_interrupted(async move { self.client @@ -276,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 @@ -291,6 +324,7 @@ impl ProducerClient { strategic_metadata: Option, chunk_size: Option, otel_trace_carrier: CarrierMap, + paused: bool, ) -> CPyResult< SubmissionId, E![ @@ -330,6 +364,7 @@ impl ProducerClient { }, metadata, strategic_metadata: strategic_metadata.unwrap_or_default(), + paused, }; self.client .insert_submission(&submission, &otel_trace_carrier) @@ -376,57 +411,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 +426,39 @@ impl ProducerClient { &self, py: Python<'_>, submission_id: SubmissionId, + timeout: Option, ) -> CPyResult< PyChunksIter, E![ FatalPythonException, + TryFromFloatSecsError, + 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) => { + 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(R(err)))), + }), + } }) }) } diff --git a/libs/opsqueue_python/tests/test_roundtrip.py b/libs/opsqueue_python/tests/test_roundtrip.py index 23d9a149..e0c95687 100644 --- a/libs/opsqueue_python/tests/test_roundtrip.py +++ b/libs/opsqueue_python/tests/test_roundtrip.py @@ -27,8 +27,11 @@ strategy_from_description, ) import logging +import time import pytest +SUBMISSION_COMPLETED_TIMEOUT = 10.0 + def increment(data: int) -> int: return data + 1 @@ -56,7 +59,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 +134,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 +153,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 +191,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 +237,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 +281,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 +323,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 +400,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 +447,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 +538,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 +574,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 +609,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 +703,93 @@ 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, + ) + + +def test_unpause_and_complete(opsqueue: OpsqueueProcess) -> None: + """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) + 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, + timeout=SUBMISSION_COMPLETED_TIMEOUT, + ) + 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/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/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/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 83e09941..fd7f67a5 100644 Binary files a/opsqueue/opsqueue_example_database_schema.db and b/opsqueue/opsqueue_example_database_schema.db differ diff --git a/opsqueue/src/common/chunk.rs b/opsqueue/src/common/chunk.rs index c27c17ec..431d0bb4 100644 --- a/opsqueue/src/common/chunk.rs +++ b/opsqueue/src/common/chunk.rs @@ -225,14 +225,13 @@ impl Chunk { #[cfg(feature = "server-logic")] pub mod db { use super::{ - Chunk, ChunkCompleted, ChunkFailed, ChunkId, ChunkIndex, ChunkSize, DateTime, SubmissionId, - Utc, u63, + 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( @@ -301,28 +300,41 @@ pub mod db { chunk_id: ChunkId, output_content: Option>, mut conn: impl WriterConnection, - ) -> Result<(), E>> { - let _chunk_size: Result>> = - conn.transaction(move |mut tx| { + ) -> Result> { + let (chunks_moved, completed_submission) = conn + .transaction(move |mut tx| { Box::pin(async move { - let completed_work = + let chunks_moved = 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()) + + let mut completed_submission = false; + if chunks_moved { + 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: \ + completed, failed, or cancelled before. Ignoring.", + chunk_id + ); + } + + Result::<(bool, bool), E>::Ok(( + chunks_moved, + completed_submission, + )) }) }) - .await; + .await?; - counter!(crate::prometheus::CHUNKS_COMPLETED_COUNTER).increment(1); - Ok(()) + if chunks_moved { + counter!(crate::prometheus::CHUNKS_COMPLETED_COUNTER).increment(1); + } + Ok(completed_submission) } /// This function MUST be called inside a transaction. @@ -335,17 +347,16 @@ 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) 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, @@ -354,26 +365,36 @@ 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. + .execute(tx.get_inner()) + .await? + .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. + // + // 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 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. + // By only updating `chunks_done` when we actually moved a chunk, we ensure that 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?; + } + Ok(chunk_moved) } /// Increment retries for a chunk, or move it to failed state. @@ -395,7 +416,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 @@ -404,23 +425,32 @@ 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) } }) @@ -583,6 +613,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, '', TRUE, 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 @@ -600,7 +717,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; @@ -620,13 +737,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 +754,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 +771,33 @@ 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")) + } + + /// 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. @@ -675,9 +818,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::*; @@ -692,11 +836,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 +896,9 @@ pub mod test { .await .expect("complete chunk failed"); - assert_eq!(count_chunks(&mut conn).await.unwrap(), u63::new(0)); - assert_eq!( - count_chunks_completed(&mut conn).await.unwrap(), - u63::new(1) - ); - assert_eq!(count_chunks_failed(&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(), 1); + assert_eq!(count_chunks_failed(&mut conn).await.unwrap(), 0); } #[sqlx::test(migrator = "crate::MIGRATOR")] @@ -770,12 +911,13 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .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 { @@ -803,6 +945,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(false)); + + let res = complete_chunk(chunk_id, None, &mut conn).await; + assert_matches!(res, Ok(false)); + } + #[sqlx::test(migrator = "crate::MIGRATOR")] pub async fn test_fail_chunk(db: sqlx::SqlitePool) { let db = WriterPool::new(db); @@ -824,11 +995,40 @@ pub mod test { .await .expect("Succeed chunk failed"); - assert_eq!(count_chunks(&mut conn).await.unwrap(), u63::new(0)); - assert_eq!( - count_chunks_completed(&mut conn).await.unwrap(), - u63::new(0) - ); - assert_eq!(count_chunks_failed(&mut conn).await.unwrap(), u63::new(1)); + assert_eq!(count_chunks(&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); + } + + #[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/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); diff --git a/opsqueue/src/common/submission.rs b/opsqueue/src/common/submission.rs index 991744dd..9c5ca320 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 { @@ -276,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::{ @@ -284,19 +309,13 @@ pub mod db { DatabaseError, E, SubmissionNotCancellable, SubmissionNotFound, TooManyMatchingSubmissions, }, + submission::SubmissionPaused, }, db::{Connection, True, WriterConnection, WriterPool}, }; 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, - SubmissionCancelled, SubmissionCompleted, SubmissionFailed, SubmissionId, SubmissionStatus, - Utc, chunk, - }; + use sqlx::{Database, QueryBuilder, Sqlite, query, query_as, query_scalar}; impl<'q> sqlx::Encode<'q, Sqlite> for SubmissionId { fn encode_by_ref( @@ -359,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(()) } @@ -430,9 +449,126 @@ 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?; + // 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(()) + }) + }) + .await + } + + #[tracing::instrument(skip(conn))] + pub(crate) async fn unpause_submission_raw( + id: SubmissionId, + mut conn: impl WriterConnection, + ) -> Result<(), E> { + 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; + ", + id, + id, + ) + .execute(conn.get_inner()) + .await?; + if res.rows_affected() == 0 { + 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 /// @@ -448,6 +584,7 @@ pub mod db { metadata: Option, strategic_metadata: StrategicMetadataMap, chunk_size: ChunkSize, + paused: bool, mut conn: impl WriterConnection, ) -> Result { let submission_id = SubmissionId::new(); @@ -463,7 +600,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)| { @@ -471,25 +608,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?!"); + } } } } @@ -580,13 +722,13 @@ 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 = $1 UNION ALL - SELECT id AS "id: SubmissionId" FROM submissions_failed 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 = $1 "#, prefix, - prefix, - prefix ) .fetch_optional(conn.get_inner()) .await?; @@ -665,8 +807,124 @@ 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 submission_row = query!( + 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()) + .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))); + } + + 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" @@ -684,23 +942,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))); - } + } + + 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, + } - let completed_row_opt = query!( + #[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" @@ -718,23 +983,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))); - } + } + + 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, + } - let failed_row_opt = query!( + #[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" @@ -754,30 +1028,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, - ))); - } + } - let cancelled_row_opt = query!( + pub(crate) struct SubmissionStatusCancelledRow { + id: SubmissionId, + prefix: Option, + chunks_total: ChunkCount, + chunks_done: ChunkCount, + metadata: Option, + strategic_metadata: sqlx::types::Json, + cancelled_at: DateTime, + } + + #[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" @@ -794,22 +1067,47 @@ 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))); - } + } - Ok(None) + 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" + , 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 + ) } #[tracing::instrument(skip(conn))] @@ -866,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:?}") } @@ -892,13 +1193,23 @@ pub mod db { /// # Errors /// /// Returns an error if cancellation or chunk skipping fails. - pub 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?; + + Ok(()) + } + Err(E::L(db_err)) => Err(E::L(db_err)), + } } #[tracing::instrument(skip(conn))] @@ -908,21 +1219,53 @@ 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()) - .await?; - if submission_opt.is_none() { + .execute(conn.get_inner()) + .await?; + if res.rows_affected() == 0 { + 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))] + pub(super) async fn cancel_paused_submission_raw( + id: SubmissionId, + mut conn: impl WriterConnection, + ) -> Result<(), E> { + let now = chrono::prelude::Utc::now(); + + 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; + ", + now, + id, + id, + ) + .execute(conn.get_inner()) + .await?; + if res.rows_affected() == 0 { Err(E::R(SubmissionNotFound(id))) } else { counter!(crate::prometheus::SUBMISSIONS_CANCELLED_COUNTER).increment(1); @@ -991,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()), @@ -1026,7 +1369,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, @@ -1048,12 +1391,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 +1408,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 +1425,50 @@ 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")) + } + + /// 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, @@ -1098,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 @@ -1106,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 ( @@ -1115,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 ( @@ -1124,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)"); @@ -1166,7 +1557,7 @@ pub mod db { Ok(()) }) }) - .await + .await } pub async fn periodically_cleanup_old(db: &WriterPool, max_age: Duration) { @@ -1194,48 +1585,46 @@ 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; + use crate::common::chunk::db::{count_chunks, count_chunks_failed, count_chunks_paused}; use crate::db::{Connection as _, WriterPool}; 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, @" @@ -1254,8 +1643,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=?) @@ -1265,114 +1653,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=?) "); } @@ -1381,7 +1701,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 +1713,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")] @@ -1435,6 +1755,7 @@ pub mod test { None, strategic_metadata.clone(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1466,15 +1787,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 +1814,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")] @@ -1522,6 +1831,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1532,6 +1842,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1542,6 +1853,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1552,6 +1864,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1584,6 +1897,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1594,6 +1908,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1604,6 +1919,7 @@ pub mod test { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -1626,18 +1942,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 @@ -1663,20 +1973,15 @@ pub mod test { StrategicMetadataMap::default(), // chunk size ChunkSize::default(), + false, &mut conn, ) .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. @@ -1774,4 +2079,133 @@ 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 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=?) + 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_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_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); + 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/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/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/server/mod.rs b/opsqueue/src/consumer/server/mod.rs index 5cd9a224..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; @@ -76,10 +77,32 @@ impl ServerState { 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); + let (completer, completer_tx) = Completer::new( + pool.writer_pool(), + &dispatcher, + config.max_chunk_retries, + status_changed_sender, + ); Self { pool, completer: Some(completer), @@ -186,6 +209,7 @@ pub struct Completer { dispatcher: Dispatcher, count: usize, max_chunk_retries: u32, + status_changed_sender: Option, } impl Completer { @@ -194,6 +218,7 @@ impl Completer { pool: &db::WriterPool, dispatcher: &Dispatcher, max_chunk_retries: u32, + status_changed_sender: Option, ) -> (Self, tokio::sync::mpsc::Sender) { let (tx, rx) = tokio::sync::mpsc::channel(1024); let pool = pool.clone(); @@ -203,6 +228,7 @@ impl Completer { dispatcher: dispatcher.clone(), count: 0, max_chunk_retries, + status_changed_sender, }; (me, tx) } @@ -237,7 +263,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 +285,9 @@ impl Completer { let _ = db::perform_explicit_wal_checkpoint(conn).await; } - db_res?; + if submission_completed? && let Some(sender) = &self.status_changed_sender { + let _ = sender.send(id.submission_id); + } Ok(()) } CompleterMessage::Fail { @@ -293,7 +321,9 @@ impl Completer { histogram!(crate::prometheus::CONSUMER_FAIL_CHUNK_DURATION) .record(start.elapsed()); - failed_permanently?; + if failed_permanently? && let Some(sender) = &self.status_changed_sender { + let _ = sender.send(id.submission_id); + } Ok(()) } } 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/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..40d851d3 --- /dev/null +++ b/opsqueue/src/delegation/server.rs @@ -0,0 +1,1157 @@ +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; +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, +) -> (Router, tokio::sync::mpsc::UnboundedSender) { + let notify_on_insert = Arc::new(Notify::new()); + let (status_changed_sender, status_changed_receiver) = tokio::sync::mpsc::unbounded_channel(); + let router = ServerState::new( + delegation_server_url, + cancellation_token.clone(), + Interface::new( + pool.clone(), + status_changed_sender.clone(), + status_changed_receiver, + ), + notify_on_insert, + ) + .run_background() + .build_router(); + + (Router::new().nest("/job", router), status_changed_sender) +} + +#[derive(Debug, Clone)] +pub struct ServerState { + interface: Interface, + cancellation_token: CancellationToken, + /// Notified when new chunks become available for dispatch (e.g. after unpausing a submission). + pub notify_on_insert: Arc, + delegation_server_url: url::Url, + http_client: reqwest::Client, +} + +impl ServerState { + pub fn new( + delegation_server_url: url::Url, + cancellation_token: CancellationToken, + interface: Interface, + notify_on_insert: Arc, + ) -> Self { + Self { + interface, + cancellation_token, + notify_on_insert, + delegation_server_url, + 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, 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.interface.pool.writer_conn().await.map_err(|e| { + tracing::error!("DB error acquiring writer connection: {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_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.interface.pool.reader_conn().await.map_err(|e| { + tracing::error!("DB error acquiring writer connection: {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 + .interface + .cancel_submissions(submission_ids) + .await + .map_err(|e| { + tracing::error!("DB error handling kill event: {e:?}"); + StatusCode::INTERNAL_SERVER_ERROR + })?; + + 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) +} + +const DELEGATION_BACKGROUND_LOOP_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5); + +async fn run_in_background( + 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, + Some(_) = state.interface.wait_for_submission_status_change() => 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 = select_out_of_date_tasks(&state.interface).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 submission status notification" + ); + } + + 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.interface.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.interface.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(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? + }; + + 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( + 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::Mutex; + use tokio::sync::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, _status_changed_sender) = app_for_tests( + &pool, + &cancellation_token, + external_server.uri().parse().unwrap(), + ); + + 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, _status_changed_sender) = app_for_tests( + &pool, + &cancellation_token, + external_server.uri().parse().unwrap(), + ); + + { + 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 (_app, status_changed_sender) = app_for_tests( + &pool, + &cancellation_token, + external_server.uri().parse().unwrap(), + ); + + 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(); + status_changed_sender.send(submission).unwrap(); + } + + 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 (_app, status_changed_sender) = app_for_tests( + &pool, + &cancellation_token, + external_server.uri().parse().unwrap(), + ); + + 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(); + status_changed_sender.send(submission).unwrap(); + } + + 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 (_app, status_changed_sender) = app_for_tests( + &pool, + &cancellation_token, + external_server.uri().parse().unwrap(), + ); + + 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(); + status_changed_sender.send(submission).unwrap(); + } + + 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 (_app, status_changed_sender) = app_for_tests( + &pool, + &cancellation_token, + external_server.uri().parse().unwrap(), + ); + + 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(); + status_changed_sender.send(submission).unwrap(); + } + + 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/client.rs b/opsqueue/src/producer/client.rs index 4f1dc92d..17a94691 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. + /// + /// 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. @@ -398,7 +441,6 @@ impl InternalProducerClientError { #[cfg(test)] #[cfg(feature = "server-logic")] mod tests { - use ux::u63; use crate::{ common::{ @@ -439,6 +481,7 @@ mod tests { None, StrategicMetadataMap::default(), ChunkSize::default(), + false, &mut conn, ) .await @@ -459,7 +502,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 { @@ -468,6 +511,7 @@ mod tests { metadata: None, strategic_metadata: StrategicMetadataMap::default(), chunk_size: None, + paused: false, }; client .insert_submission(&submission, &std::collections::HashMap::default()) @@ -477,7 +521,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 +539,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")] @@ -511,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()) @@ -525,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) => { @@ -535,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 74b87075..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}; @@ -27,6 +28,7 @@ pub struct ServerState { pool: DBPools, notify_on_insert: Arc, max_submissions: MaxSubmissions, + status_changed_sender: Option, } impl ServerState { @@ -34,11 +36,21 @@ impl ServerState { pool: DBPools, notify_on_insert: 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, max_submissions, + status_changed_sender, } } @@ -67,6 +79,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), @@ -127,7 +143,12 @@ 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(()) => { + 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()), Err(R(L(not_found_err))) => { Err((StatusCode::NOT_FOUND, Json(not_found_err)).into_response()) @@ -138,6 +159,32 @@ 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(); + 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()), + 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 +246,7 @@ async fn insert_submission( request.metadata, request.strategic_metadata, request.chunk_size.unwrap_or_default(), + request.paused, &mut conn, ) .await?; @@ -208,8 +256,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)) } @@ -222,7 +272,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 +280,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..99e76439 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" + ); describe_histogram!( SUBMISSIONS_DURATION_COMPLETE_HISTOGRAM, Unit::Seconds, @@ -211,9 +223,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?; diff --git a/opsqueue/src/server.rs b/opsqueue/src/server.rs index 5f9ec428..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,34 +97,55 @@ 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 (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(), cancellation_token.clone(), reservation_expiration, config, + Some(status_changed_sender.clone()), ) .run_background() .build_router(); - let producer_routes = crate::producer::server::ServerState::new( - pool, - notify_on_insert, + let producer_routes = crate::producer::server::ServerState::new_with_status_sender( + pool.clone(), + notify_on_insert.clone(), config.max_submissions_returned, + Some(status_changed_sender), ) .build_router(); - let routes = Router::new() + let mut routes = Router::new() .nest("/producer", producer_routes) .nest("/consumer", consumer_routes); + if let Some(delegation_server_url) = config.delegation_server_url.clone() { + let delegation_routes = crate::delegation::server::ServerState::new( + delegation_server_url, + cancellation_token.clone(), + interface, + notify_on_insert.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/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); + } + } +} 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]