Skip to main content

vortex_layout/layouts/
buffered.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright the Vortex contributors
3
4use std::collections::VecDeque;
5use std::sync::Arc;
6
7use async_stream::try_stream;
8use async_trait::async_trait;
9use futures::StreamExt as _;
10use futures::pin_mut;
11use vortex_error::VortexResult;
12use vortex_session::VortexSession;
13
14use crate::LayoutRef;
15use crate::LayoutStrategy;
16use crate::LayoutWriterContext;
17use crate::segments::SegmentSinkRef;
18use crate::sequence::SendableSequentialStream;
19use crate::sequence::SequencePointer;
20use crate::sequence::SequentialStreamAdapter;
21use crate::sequence::SequentialStreamExt as _;
22
23#[derive(Clone)]
24pub struct BufferedStrategy {
25    child: Arc<dyn LayoutStrategy>,
26    buffer_size: u64,
27}
28
29impl BufferedStrategy {
30    pub fn new<S: LayoutStrategy>(child: S, buffer_size: u64) -> Self {
31        Self {
32            child: Arc::new(child),
33            buffer_size,
34        }
35    }
36}
37
38#[async_trait]
39impl LayoutStrategy for BufferedStrategy {
40    async fn write_stream(
41        &self,
42        ctx: LayoutWriterContext,
43        segment_sink: SegmentSinkRef,
44        stream: SendableSequentialStream,
45        eof: SequencePointer,
46        session: &VortexSession,
47    ) -> VortexResult<LayoutRef> {
48        let dtype = stream.dtype().clone();
49        let buffer_size = self.buffer_size;
50        let buffered_bytes = ctx.buffered_bytes_tracker().clone();
51
52        let buffered_stream = try_stream! {
53            let stream = stream.peekable();
54            pin_mut!(stream);
55
56            let mut nbytes = 0u64;
57            let mut chunks = VecDeque::new();
58
59            while let Some(chunk) = stream.as_mut().next().await {
60                let (sequence_id, chunk) = chunk?;
61                let chunk_size = chunk.nbytes();
62                nbytes += chunk_size;
63                chunks.push_back((chunk, buffered_bytes.reserve(chunk_size)));
64
65                // If this is the last element, flush everything.
66                if stream.as_mut().peek().await.is_none() {
67                    let mut sequence_ptr = sequence_id.descend();
68                    while let Some((chunk, reservation)) = chunks.pop_front() {
69                        drop(reservation);
70                        yield (sequence_ptr.advance(), chunk)
71                    }
72                    break;
73                }
74
75                if nbytes < 2 * buffer_size {
76                    continue;
77                };
78
79                // Wait until we're at 2x the buffer size before flushing 1x the buffer size.
80                // This avoids small tail stragglers being flushed at the end of the file.
81                let mut sequence_ptr = sequence_id.descend();
82                while nbytes > buffer_size {
83                    let Some((chunk, reservation)) = chunks.pop_front() else {
84                        break;
85                    };
86                    nbytes -= reservation.bytes();
87                    drop(reservation);
88                    yield (sequence_ptr.advance(), chunk)
89                }
90            }
91        };
92
93        self.child
94            .write_stream(
95                ctx,
96                segment_sink,
97                SequentialStreamAdapter::new(dtype, buffered_stream).sendable(),
98                eof,
99                session,
100            )
101            .await
102    }
103}