use crate::compression::Compression;
use crate::enums::IPCMessageProtocol;
use crate::models::codecs::ipc::ArrowIpcCodec;
use crate::models::writers::ipc::table_stream::TableStreamWriter;
use crate::traits::stream_buffer::StreamBuffer;
use minarrow::{Field, TableV, Vec64};
use std::io;
use tokio::io::AsyncWrite;
use futures_sink::Sink;
use std::pin::Pin;
use std::task::{Context, Poll};
pub type TableSink<W> = GTableSink<W, Vec<u8>>;
pub type TableSink64<W> = GTableSink<W, Vec64<u8>>;
pub struct GTableSink<W, B>
where
W: AsyncWrite + Unpin + Send + Sync + 'static,
B: StreamBuffer + Unpin + 'static,
{
pub(crate) schema: Vec<Field>,
pub(crate) codec: ArrowIpcCodec<B>,
pub(crate) destination: W,
pub(crate) protocol: IPCMessageProtocol,
pub(crate) finished: bool,
pub(crate) frame_buf: Option<B>, pub(crate) frame_pos: usize, pub(crate) encode_buf: B,
pub(crate) file_writer: Option<TableStreamWriter<B>>,
}
impl<W, B> GTableSink<W, B>
where
W: AsyncWrite + Unpin + Send + Sync + 'static,
B: StreamBuffer + std::fmt::Debug + Unpin + 'static,
{
pub fn new(
sink: W,
schema: Vec<Field>,
protocol: IPCMessageProtocol,
compression: Option<Compression>,
) -> io::Result<Self> {
let file_writer = if protocol == IPCMessageProtocol::File {
Some(TableStreamWriter::new(schema.clone(), protocol, compression))
} else {
None
};
Ok(Self {
codec: ArrowIpcCodec::new(schema.clone(), protocol, compression, None),
schema,
destination: sink,
protocol,
finished: false,
frame_buf: None,
frame_pos: 0,
encode_buf: B::with_capacity(0),
file_writer,
})
}
pub fn sink_mut(&mut self) -> &mut W {
&mut self.destination
}
pub(crate) fn encode_frame(
&mut self,
view: TableV,
custom_metadata: Option<&[(String, String)]>,
) -> io::Result<()> {
if self.protocol == IPCMessageProtocol::Stream {
let mut buf = std::mem::replace(&mut self.encode_buf, B::with_capacity(0));
let len = buf.len();
if len > 0 {
buf.drain(0..len);
}
self.codec
.encode_stream_batch(&view, &mut buf, 0, custom_metadata)?;
self.frame_buf = Some(buf);
self.frame_pos = 0;
} else if let Some(writer) = &mut self.file_writer {
writer.write(view)?;
}
Ok(())
}
}
impl<W, B> Sink<TableV> for GTableSink<W, B>
where
W: AsyncWrite + Unpin + Send + Sync + 'static,
B: StreamBuffer + std::fmt::Debug + Unpin + 'static,
{
type Error = io::Error;
fn poll_ready(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
if self.frame_buf.is_some() {
return self.as_mut().poll_flush(cx);
}
Poll::Ready(Ok(()))
}
fn start_send(self: Pin<&mut Self>, view: TableV) -> Result<(), Self::Error> {
self.get_mut().encode_frame(view, None)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
loop {
if self.frame_buf.is_none() {
if let Some(writer) = &mut self.file_writer {
if let Some(Ok(frame)) = writer.next_frame() {
self.frame_pos = 0;
self.frame_buf = Some(frame);
} else {
break;
}
} else {
break;
}
}
if let Some(buf) = self.frame_buf.take() {
let remaining = &buf.as_ref()[self.frame_pos..];
const MAX_WRITE_CHUNK: usize = 1024 * 1024; let chunk = if remaining.len() > MAX_WRITE_CHUNK {
&remaining[..MAX_WRITE_CHUNK]
} else {
remaining
};
match Pin::new(&mut self.destination).poll_write(cx, chunk) {
Poll::Pending => {
self.frame_buf = Some(buf);
return Poll::Pending;
}
Poll::Ready(Ok(0)) => return Poll::Ready(Err(io::ErrorKind::WriteZero.into())),
Poll::Ready(Ok(n)) => {
self.frame_pos += n;
if self.frame_pos < buf.as_ref().len() {
self.frame_buf = Some(buf);
cx.waker().wake_by_ref();
return Poll::Pending;
} else {
self.encode_buf = buf;
self.frame_pos = 0;
}
}
Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
}
} else {
break; }
}
Pin::new(&mut self.destination).poll_flush(cx)
}
fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
if !self.finished {
if let Some(writer) = &mut self.file_writer {
writer.finish()?;
} else {
let mut eos_buf = B::with_capacity(8);
self.codec.finish(&mut eos_buf)?;
self.frame_buf = Some(eos_buf);
self.frame_pos = 0;
}
self.finished = true;
}
match self.as_mut().poll_flush(cx)? {
Poll::Pending => return Poll::Pending,
Poll::Ready(()) => { }
}
Pin::new(&mut self.destination)
.poll_shutdown(cx)
.map_err(Into::into)
}
}
#[cfg(test)]
mod tests {
use futures_util::SinkExt;
use minarrow::Table;
use super::*;
use crate::models::decoders::ipc::table_stream_decoder::TableStreamDecoder;
use crate::test_helpers::{int32_col, int64_col};
fn test_table() -> Table {
Table::new("t".to_string(), vec![int32_col(), int64_col()].into())
}
fn schema_of(table: &Table) -> Vec<Field> {
table
.cols
.iter()
.map(|fa| fa.field.as_ref().clone())
.collect()
}
async fn decoded_batches(bytes: Vec<u8>) -> usize {
let mut decoder = TableStreamDecoder::<Vec64<u8>>::new(
std::io::Cursor::new(bytes),
64 * 1024,
IPCMessageProtocol::Stream,
None,
);
let mut count = 0;
while let Some(item) = decoder.read_keyed().await {
item.expect("stream decoded cleanly");
count += 1;
}
count
}
#[tokio::test]
async fn feed_without_flush_keeps_every_table() {
let table = test_table();
let mut sink =
TableSink64::new(Vec::new(), schema_of(&table), IPCMessageProtocol::Stream, None)
.unwrap();
sink.feed(table.clone().into()).await.unwrap();
sink.feed(table.clone().into()).await.unwrap();
sink.feed(table.clone().into()).await.unwrap();
SinkExt::flush(&mut sink).await.unwrap();
SinkExt::close(&mut sink).await.unwrap();
let bytes = std::mem::take(&mut sink.destination);
assert_eq!(decoded_batches(bytes).await, 3);
}
#[tokio::test]
async fn send_per_item_keeps_every_table() {
let table = test_table();
let mut sink =
TableSink64::new(Vec::new(), schema_of(&table), IPCMessageProtocol::Stream, None)
.unwrap();
SinkExt::send(&mut sink, table.clone().into()).await.unwrap();
SinkExt::send(&mut sink, table.clone().into()).await.unwrap();
SinkExt::close(&mut sink).await.unwrap();
let bytes = std::mem::take(&mut sink.destination);
assert_eq!(decoded_batches(bytes).await, 2);
}
}