Skip to content

Commit ba2813c

Browse files
committed
Feat: stream completed chunks
1 parent c239da7 commit ba2813c

5 files changed

Lines changed: 325 additions & 7 deletions

File tree

‎.vscode/opsqueue.code-workspace‎

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,13 @@
1+
{
2+
"folders": [
3+
{
4+
"path": "..",
5+
"name": "opsqueue",
6+
},
7+
{
8+
"path": "../libs/opsqueue_python",
9+
"name": "opsqueue_python",
10+
}
11+
],
12+
"settings": {}
13+
}

‎libs/opsqueue_python/python/opsqueue/producer.py‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -231,6 +231,14 @@ def run_submission_chunks(
231231
)
232232
return self.blocking_stream_completed_submission_chunks(submission_id)
233233

234+
def stream_submission_chunks(self, submission_id: SubmissionId) -> Iterator[bytes]:
235+
return self.inner.stream_submission_chunks(submission_id) # type: ignore[no-any-return]
236+
237+
async def async_stream_submission_chunks(
238+
self, submission_id: SubmissionId
239+
) -> AsyncIterator[bytes]:
240+
return await self.inner.async_stream_submission_chunks(submission_id) # type: ignore[no-any-return]
241+
234242
async def async_run_submission_chunks(
235243
self,
236244
chunk_contents: Iterable[bytes],

‎libs/opsqueue_python/src/errors.rs‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -153,6 +153,7 @@ impl From<CError<TooManyMatchingSubmissions>> for PyErr {
153153
}
154154
}
155155

156+
#[derive(Debug)]
156157
pub struct SubmissionFailed(
157158
pub crate::common::SubmissionFailed,
158159
pub crate::common::ChunkFailed,

‎libs/opsqueue_python/src/producer.rs‎

Lines changed: 209 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -427,6 +427,146 @@ impl ProducerClient {
427427
Ok(res)
428428
}
429429

430+
/// Stream output chunks as soon as each consumer has completed them.
431+
pub fn stream_submission_chunks(&self, submission_id: SubmissionId) -> PyChunksIter {
432+
self.streaming_submission_chunks(submission_id)
433+
}
434+
435+
fn streaming_submission_chunks(&self, submission_id: SubmissionId) -> PyChunksIter {
436+
let client = self.client.clone();
437+
let object_store_client = self.object_store_client.clone();
438+
let stream = futures::stream::unfold(
439+
(
440+
client,
441+
object_store_client,
442+
submission_id,
443+
u63::new(0),
444+
None,
445+
Duration::from_millis(10),
446+
),
447+
|(client, object_store_client, submission_id, index, prefix, interval)| async move {
448+
let mut interval = interval;
449+
loop {
450+
let status = match client.get_submission(submission_id.into()).await {
451+
Ok(Some(status)) => status,
452+
Ok(None) => {
453+
return Some((
454+
Err(StreamingChunkError::SubmissionNotFound),
455+
(
456+
client,
457+
object_store_client,
458+
submission_id,
459+
index,
460+
prefix,
461+
interval,
462+
),
463+
));
464+
}
465+
Err(error) => {
466+
return Some((
467+
Err(StreamingChunkError::Internal(error)),
468+
(
469+
client,
470+
object_store_client,
471+
submission_id,
472+
index,
473+
prefix,
474+
interval,
475+
),
476+
));
477+
}
478+
};
479+
480+
match status {
481+
submission::SubmissionStatus::InProgress(submission) => {
482+
let prefix = prefix.clone().or(submission.prefix);
483+
if index < submission.chunks_done.into() {
484+
let prefix = prefix
485+
.expect("in-progress submissions have an object-store prefix");
486+
let result = object_store_client
487+
.retrieve_chunk(&prefix, index.into(), ChunkType::Output)
488+
.await
489+
.map_err(StreamingChunkError::Retrieval);
490+
return Some((
491+
result,
492+
(
493+
client,
494+
object_store_client,
495+
submission_id,
496+
index + u63::new(1),
497+
Some(prefix),
498+
interval,
499+
),
500+
));
501+
}
502+
}
503+
submission::SubmissionStatus::Completed(submission) => {
504+
let prefix = prefix.clone().or(submission.prefix);
505+
if index < submission.chunks_total.into() {
506+
let prefix = prefix
507+
.expect("completed submissions have an object-store prefix");
508+
let result = object_store_client
509+
.retrieve_chunk(&prefix, index.into(), ChunkType::Output)
510+
.await
511+
.map_err(StreamingChunkError::Retrieval);
512+
return Some((
513+
result,
514+
(
515+
client,
516+
object_store_client,
517+
submission_id,
518+
index + u63::new(1),
519+
Some(prefix),
520+
interval,
521+
),
522+
));
523+
}
524+
return None;
525+
}
526+
submission::SubmissionStatus::Failed(submission, chunk) => {
527+
let failure =
528+
crate::common::ChunkFailed::from_internal(chunk, &submission);
529+
return Some((
530+
Err(StreamingChunkError::Failed(
531+
crate::errors::SubmissionFailed(submission.into(), failure),
532+
)),
533+
(
534+
client,
535+
object_store_client,
536+
submission_id,
537+
index,
538+
prefix,
539+
interval,
540+
),
541+
));
542+
}
543+
submission::SubmissionStatus::Cancelled(_) => {
544+
return Some((
545+
Err(StreamingChunkError::Cancelled),
546+
(
547+
client,
548+
object_store_client,
549+
submission_id,
550+
index,
551+
prefix,
552+
interval,
553+
),
554+
));
555+
}
556+
}
557+
558+
tokio::time::sleep(interval).await;
559+
if interval < SUBMISSION_POLLING_INTERVAL {
560+
interval = (interval * 2).min(SUBMISSION_POLLING_INTERVAL);
561+
}
562+
}
563+
},
564+
)
565+
.map(|item| item.map_err(CError))
566+
.boxed();
567+
PyChunksIter::from_stream(self, stream)
568+
}
569+
430570
/// Blocks (and short-polls) until the submission is completed.
431571
///
432572
/// We start with a small short-polling interval
@@ -457,6 +597,30 @@ impl ProducerClient {
457597
})
458598
}
459599

600+
/// Return an awaitable that resolves immediately to an async iterator of output chunks.
601+
///
602+
/// The iterator polls submission progress and yields each output chunk as soon as it is ready.
603+
///
604+
/// # Errors
605+
///
606+
/// Returns a Python error if creating the awaitable fails.
607+
pub fn async_stream_submission_chunks<'p>(
608+
&self,
609+
py: Python<'p>,
610+
submission_id: SubmissionId,
611+
) -> PyResult<Bound<'p, PyAny>> {
612+
let me = self.clone();
613+
let _tokio_active_runtime_guard = me.runtime.enter();
614+
async_util::future_into_py(
615+
py,
616+
async_util::async_detach(Box::pin(async move {
617+
Ok(PyChunksAsyncIter::from(
618+
me.streaming_submission_chunks(submission_id),
619+
))
620+
})),
621+
)
622+
}
623+
460624
/// Return an awaitable that resolves to an async iterator of output chunks.
461625
///
462626
/// # Errors
@@ -565,7 +729,41 @@ impl ProducerClient {
565729
}
566730
}
567731

568-
pub type ChunksStream = BoxStream<'static, CPyResult<Vec<u8>, ChunkRetrievalError>>;
732+
#[derive(Debug)]
733+
enum StreamingChunkError {
734+
Retrieval(ChunkRetrievalError),
735+
Internal(InternalProducerClientError),
736+
Failed(crate::errors::SubmissionFailed),
737+
SubmissionNotFound,
738+
Cancelled,
739+
}
740+
741+
impl std::fmt::Display for StreamingChunkError {
742+
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
743+
match self {
744+
Self::Retrieval(error) => error.fmt(f),
745+
Self::Internal(error) => error.fmt(f),
746+
Self::Failed(_) => write!(f, "Submission failed"),
747+
Self::SubmissionNotFound => write!(f, "Submission not found"),
748+
Self::Cancelled => write!(f, "Submission cancelled"),
749+
}
750+
}
751+
}
752+
753+
impl std::error::Error for StreamingChunkError {}
754+
755+
impl From<CError<StreamingChunkError>> for PyErr {
756+
fn from(value: CError<StreamingChunkError>) -> Self {
757+
match value.0 {
758+
StreamingChunkError::Retrieval(error) => CError(error).into(),
759+
StreamingChunkError::Internal(error) => CError(error).into(),
760+
StreamingChunkError::Failed(error) => CError(error).into(),
761+
error => PyException::new_err(error.to_string()),
762+
}
763+
}
764+
}
765+
766+
type ChunksStream = BoxStream<'static, CPyResult<Vec<u8>, StreamingChunkError>>;
569767

570768
#[pyclass(module = "opsqueue")]
571769
pub struct PyChunksIter {
@@ -574,16 +772,20 @@ pub struct PyChunksIter {
574772
}
575773

576774
impl PyChunksIter {
775+
fn from_stream(client: &ProducerClient, stream: ChunksStream) -> Self {
776+
Self {
777+
stream: Arc::new(tokio::sync::Mutex::new(stream)),
778+
runtime: client.runtime.clone(),
779+
}
780+
}
781+
577782
pub(crate) fn new(client: &ProducerClient, prefix: String, chunks_total: u63) -> Self {
578783
let stream = client
579784
.object_store_client
580785
.retrieve_chunks(prefix, chunks_total, ChunkType::Output)
581-
.map_err(CError)
786+
.map_err(|error| CError(StreamingChunkError::Retrieval(error)))
582787
.boxed();
583-
Self {
584-
stream: Arc::new(tokio::sync::Mutex::new(stream)),
585-
runtime: client.runtime.clone(),
586-
}
788+
Self::from_stream(client, stream)
587789
}
588790
}
589791

@@ -593,7 +795,7 @@ impl PyChunksIter {
593795
slf
594796
}
595797

596-
fn __next__(&self, py: Python<'_>) -> Option<CPyResult<Vec<u8>, ChunkRetrievalError>> {
798+
fn __next__(&self, py: Python<'_>) -> Option<CPyResult<Vec<u8>, StreamingChunkError>> {
597799
// The only time we need the GIL is when turning the result back.
598800
// By unlocking here, we reduce the chance of deadlocks.
599801
py.detach(move || {

‎libs/opsqueue_python/tests/test_roundtrip.py‎

Lines changed: 94 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,8 @@
33
# - use `RUST_LOG="opsqueue=info"` (or `opsqueue=debug` or `debug` for even more verbosity), together with to the pytest option `-s` AKA `--capture=no`, to debug the opsqueue binary itself.
44

55
from collections.abc import Iterator, Sequence
6+
import asyncio
7+
import time
68
from opsqueue.producer import (
79
SubmissionId,
810
ProducerClient,
@@ -667,3 +669,95 @@ def test_lookup_too_many_submission_ids_by_strategic_metadata() -> None:
667669
)
668670
assert exc.type is TooManyMatchingSubmissionsError
669671
assert exc.value.max_submissions == max_
672+
673+
674+
def test_streams_completed_chunks_before_submission_finishes(
675+
opsqueue: OpsqueueProcess,
676+
any_consumer_strategy: StrategyDescription,
677+
) -> None:
678+
url = "file:///tmp/opsqueue/test_streaming_results"
679+
producer_client = ProducerClient(f"localhost:{opsqueue.port}", url)
680+
submission_id = producer_client.insert_submission_chunks(
681+
[b"[1]", b"[2]"], chunk_size=1
682+
)
683+
684+
def complete_chunks(
685+
_submission_id_value: int,
686+
strategy: StrategyDescription,
687+
) -> None:
688+
consumer_client = ConsumerClient(f"localhost:{opsqueue.port}", url)
689+
chunks = sorted(
690+
consumer_client.reserve_chunks(
691+
max=2,
692+
strategy=strategy_from_description(strategy),
693+
),
694+
key=lambda chunk: chunk.chunk_index,
695+
)
696+
consumer_client.complete_chunk(
697+
chunks[0].submission_id,
698+
chunks[0].submission_prefix,
699+
chunks[0].chunk_index,
700+
chunks[0].input_content,
701+
)
702+
time.sleep(0.25)
703+
consumer_client.complete_chunk(
704+
chunks[1].submission_id,
705+
chunks[1].submission_prefix,
706+
chunks[1].chunk_index,
707+
chunks[1].input_content,
708+
)
709+
710+
with background_process(
711+
complete_chunks,
712+
args=(submission_id.id, any_consumer_strategy),
713+
):
714+
results = producer_client.stream_submission_chunks(submission_id)
715+
assert next(results) == b"[1]"
716+
assert next(results) == b"[2]"
717+
718+
719+
def test_async_streams_completed_chunks_before_submission_finishes(
720+
opsqueue: OpsqueueProcess,
721+
any_consumer_strategy: StrategyDescription,
722+
) -> None:
723+
url = "file:///tmp/opsqueue/test_async_streaming_results"
724+
producer_client = ProducerClient(f"localhost:{opsqueue.port}", url)
725+
submission_id = producer_client.insert_submission_chunks(
726+
[b"[1]", b"[2]"], chunk_size=1
727+
)
728+
729+
def complete_chunks(
730+
_submission_id_value: int,
731+
strategy: StrategyDescription,
732+
) -> None:
733+
consumer_client = ConsumerClient(f"localhost:{opsqueue.port}", url)
734+
chunks = sorted(
735+
consumer_client.reserve_chunks(
736+
max=2,
737+
strategy=strategy_from_description(strategy),
738+
),
739+
key=lambda chunk: chunk.chunk_index,
740+
)
741+
consumer_client.complete_chunk(
742+
chunks[0].submission_id,
743+
chunks[0].submission_prefix,
744+
chunks[0].chunk_index,
745+
chunks[0].input_content,
746+
)
747+
time.sleep(0.25)
748+
consumer_client.complete_chunk(
749+
chunks[1].submission_id,
750+
chunks[1].submission_prefix,
751+
chunks[1].chunk_index,
752+
chunks[1].input_content,
753+
)
754+
755+
async def collect() -> list[bytes]:
756+
results = await producer_client.async_stream_submission_chunks(submission_id)
757+
return [chunk async for chunk in results]
758+
759+
with background_process(
760+
complete_chunks,
761+
args=(submission_id.id, any_consumer_strategy),
762+
):
763+
assert asyncio.run(collect()) == [b"[1]", b"[2]"]

0 commit comments

Comments
 (0)