use std::sync::Arc;
use std::sync::atomic::AtomicU64;
use std::sync::atomic::Ordering;
use async_trait::async_trait;
use futures::StreamExt;
use vortex_array::ArrayContext;
use vortex_array::ArrayId;
use vortex_array::normalize::NormalizeOptions;
use vortex_array::normalize::Operation;
use vortex_error::VortexResult;
use vortex_session::VortexSession;
use vortex_utils::aliases::hash_set::HashSet;
use crate::LayoutRef;
use crate::segments::SegmentSinkRef;
use crate::sequence::SendableSequentialStream;
use crate::sequence::SequencePointer;
use crate::sequence::SequentialStreamAdapter;
use crate::sequence::SequentialStreamExt;
#[derive(Clone, Debug, Default)]
pub struct BufferedBytesTracker(Arc<AtomicU64>);
impl BufferedBytesTracker {
pub fn new() -> Self {
Self::default()
}
pub fn buffered_bytes(&self) -> u64 {
self.0.load(Ordering::Relaxed)
}
pub fn reserve(&self, bytes: u64) -> BufferedBytesReservation {
self.0.fetch_add(bytes, Ordering::Relaxed);
BufferedBytesReservation {
tracker: self.clone(),
bytes,
}
}
}
#[derive(Debug)]
pub struct BufferedBytesReservation {
tracker: BufferedBytesTracker,
bytes: u64,
}
impl BufferedBytesReservation {
pub fn bytes(&self) -> u64 {
self.bytes
}
}
impl Drop for BufferedBytesReservation {
fn drop(&mut self) {
self.tracker.0.fetch_sub(self.bytes, Ordering::Relaxed);
}
}
#[derive(Clone)]
pub struct LayoutWriterContext {
array_ctx: ArrayContext,
buffered_bytes: BufferedBytesTracker,
}
impl LayoutWriterContext {
pub fn new(array_ctx: ArrayContext) -> Self {
Self {
array_ctx,
buffered_bytes: BufferedBytesTracker::new(),
}
}
pub fn with_buffered_bytes_tracker(mut self, tracker: BufferedBytesTracker) -> Self {
self.buffered_bytes = tracker;
self
}
pub fn array_ctx(&self) -> &ArrayContext {
&self.array_ctx
}
pub fn buffered_bytes_tracker(&self) -> &BufferedBytesTracker {
&self.buffered_bytes
}
pub fn buffered_bytes(&self) -> u64 {
self.buffered_bytes.buffered_bytes()
}
pub fn reserve_buffered_bytes(&self, bytes: u64) -> BufferedBytesReservation {
self.buffered_bytes.reserve(bytes)
}
}
impl From<ArrayContext> for LayoutWriterContext {
fn from(array_ctx: ArrayContext) -> Self {
Self::new(array_ctx)
}
}
#[async_trait]
pub trait LayoutStrategy: 'static + Send + Sync {
async fn write_stream(
&self,
ctx: LayoutWriterContext,
segment_sink: SegmentSinkRef,
stream: SendableSequentialStream,
eof: SequencePointer,
session: &VortexSession,
) -> VortexResult<LayoutRef>;
}
#[derive(Clone)]
pub struct LayoutStrategyEncodingValidator {
child: Arc<dyn LayoutStrategy>,
allowed_encodings: Arc<HashSet<ArrayId>>,
}
impl LayoutStrategyEncodingValidator {
pub fn new<S: LayoutStrategy>(child: S, allowed_encodings: HashSet<ArrayId>) -> Self {
Self {
child: Arc::new(child),
allowed_encodings: Arc::new(allowed_encodings),
}
}
}
#[async_trait]
impl LayoutStrategy for LayoutStrategyEncodingValidator {
async fn write_stream(
&self,
ctx: LayoutWriterContext,
segment_sink: SegmentSinkRef,
stream: SendableSequentialStream,
eof: SequencePointer,
session: &VortexSession,
) -> VortexResult<LayoutRef> {
let dtype = stream.dtype().clone();
let allowed_encodings = Arc::clone(&self.allowed_encodings);
let stream = stream.map(move |chunk| {
let (sequence_id, chunk) = chunk?;
let chunk = chunk.normalize(&mut NormalizeOptions {
allowed: &allowed_encodings,
operation: Operation::Error,
})?;
Ok((sequence_id, chunk))
});
self.child
.write_stream(
ctx,
segment_sink,
SequentialStreamAdapter::new(dtype, stream).sendable(),
eof,
session,
)
.await
}
}
#[async_trait]
impl LayoutStrategy for Arc<dyn LayoutStrategy> {
async fn write_stream(
&self,
ctx: LayoutWriterContext,
segment_sink: SegmentSinkRef,
stream: SendableSequentialStream,
eof: SequencePointer,
session: &VortexSession,
) -> VortexResult<LayoutRef> {
(**self)
.write_stream(ctx, segment_sink, stream, eof, session)
.await
}
}
#[cfg(test)]
mod tests {
use crate::strategy::BufferedBytesTracker;
#[test]
fn reservations_accumulate_and_release() {
let tracker = BufferedBytesTracker::new();
assert_eq!(tracker.buffered_bytes(), 0);
let first = tracker.reserve(16);
let second = tracker.reserve(32);
assert_eq!(tracker.buffered_bytes(), 48);
assert_eq!(first.bytes(), 16);
drop(first);
assert_eq!(tracker.buffered_bytes(), 32);
drop(second);
assert_eq!(tracker.buffered_bytes(), 0);
}
#[test]
fn clones_share_the_same_counter() {
let tracker = BufferedBytesTracker::new();
let observer = tracker.clone();
let reservation = tracker.reserve(8);
assert_eq!(observer.buffered_bytes(), 8);
drop(reservation);
assert_eq!(observer.buffered_bytes(), 0);
}
}