use std::sync::Arc;
use std::sync::atomic::AtomicU64;
use std::sync::atomic::Ordering;
use async_trait::async_trait;
use vortex_array::ArrayContext;
use vortex_array::aggregate_fn::AggregateFnId;
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;
#[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,
allowed_aggregates: Option<Arc<HashSet<AggregateFnId>>>,
buffered_bytes: BufferedBytesTracker,
}
impl LayoutWriterContext {
pub fn new(array_ctx: ArrayContext) -> Self {
Self {
array_ctx,
allowed_aggregates: None,
buffered_bytes: BufferedBytesTracker::new(),
}
}
pub fn with_allowed_aggregates(mut self, allowed: HashSet<AggregateFnId>) -> Self {
self.allowed_aggregates = Some(Arc::new(allowed));
self
}
pub fn allows_aggregate(&self, aggregate: &AggregateFnId) -> bool {
self.allowed_aggregates
.as_ref()
.is_none_or(|allowed| allowed.contains(aggregate))
}
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>;
}
#[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);
}
}