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(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
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)
}
}