use std::collections::HashMap;
use std::io;
use minarrow::structs::shared_buffer::SharedBuffer;
use minarrow::{Field, Table, Vec64};
use crate::compression::Compression;
use crate::enums::IPCMessageProtocol;
use crate::models::decoders::ipc::{decode_ipc_frame, decode_ipc_payload};
use crate::models::decoders::limits::DecodeLimits;
use crate::models::encoders::ipc::record_batch::encode_record_batch;
use crate::models::encoders::ipc::table_stream::TableStreamEncoder;
use crate::traits::decoder::Decoder;
use crate::traits::encoder::Encoder;
use crate::traits::stream_buffer::StreamBuffer;
pub use crate::models::frames::ipc_message::IPCFrameResult;
pub struct ArrowIpcCodec<B: StreamBuffer> {
pub(crate) encoder: TableStreamEncoder<B>,
fields: Vec<Field>,
dicts: HashMap<i64, Vec<String>>,
shared_cache: Option<SharedBuffer>,
limits: DecodeLimits,
}
impl<B: StreamBuffer + Unpin> ArrowIpcCodec<B> {
pub fn new(
schema: Vec<Field>,
protocol: IPCMessageProtocol,
compression: Option<Compression>,
limits: Option<DecodeLimits>,
) -> Self {
Self {
encoder: TableStreamEncoder::new(schema, protocol, compression),
fields: Vec::new(),
dicts: HashMap::new(),
shared_cache: None,
limits: limits.unwrap_or_default(),
}
}
pub fn limits(&self) -> DecodeLimits {
self.limits
}
pub fn encode_stream_batch(
&mut self,
view: &minarrow::TableV,
out: &mut B,
base_offset: usize,
custom_metadata: Option<&[(String, String)]>,
) -> io::Result<usize> {
encode_record_batch(&mut self.encoder, view, out, base_offset, custom_metadata)
}
pub fn decode_payload(&mut self, payload: SharedBuffer) -> io::Result<minarrow::Table> {
let (table, shared) = decode_ipc_payload::<B>(
payload,
&mut self.fields,
&mut self.dicts,
self.shared_cache.take(),
self.limits,
)?;
self.shared_cache = Some(shared);
Ok(table)
}
pub fn decode_stream(&mut self, bytes: Vec64<u8>) -> io::Result<Vec<minarrow::Table>>
where
B: 'static,
{
use crate::enums::DecodeResult;
use crate::models::decoders::ipc::ArrowIPCFrameDecoder;
use crate::models::frames::ipc_message::IPCFrameResult;
use crate::traits::frame_decoder::FrameDecoder;
let shared = SharedBuffer::from_vec64(bytes);
let mut frame_decoder: ArrowIPCFrameDecoder<B> =
ArrowIPCFrameDecoder::new(self.encoder.protocol, Some(self.limits));
let total = shared.len();
let mut tables: Vec<minarrow::Table> = Vec::new();
let mut pos = 0;
while pos < total {
let buf = &shared.as_slice()[pos..];
match frame_decoder.decode(buf)? {
DecodeResult::Frame { frame, consumed } => {
let meta_range = frame.message_range;
let body_range = frame.body_range;
if meta_range.is_empty() && body_range.is_empty() {
pos += consumed;
if pos < total {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"{} bytes of trailing data after Arrow IPC EOS marker",
total - pos
),
));
}
break;
}
let meta_bytes =
&shared.as_slice()[pos + meta_range.start..pos + meta_range.end];
let body_len = body_range.end - body_range.start;
let body_shared = shared.slice(pos + body_range.start..pos + body_range.end);
match self.decode_frame(meta_bytes, body_shared, body_len)? {
IPCFrameResult::Batch(table) => tables.push(table),
IPCFrameResult::Schema
| IPCFrameResult::Dictionary
| IPCFrameResult::EndOfStream => {}
}
pos += consumed;
}
DecodeResult::NeedMore => {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"Incomplete IPC frame in multi-batch buffer",
));
}
}
}
Ok(tables)
}
pub fn decode_frame(
&mut self,
message: &[u8],
body: SharedBuffer,
body_len: usize,
) -> io::Result<IPCFrameResult> {
decode_ipc_frame(
message,
body,
body_len,
&mut self.fields,
&mut self.dicts,
&mut self.shared_cache,
self.limits,
)
}
pub fn schema(&self) -> &[Field] {
&self.fields
}
pub fn dicts(&self) -> &HashMap<i64, Vec<String>> {
&self.dicts
}
pub fn protocol(&self) -> IPCMessageProtocol {
self.encoder.protocol
}
pub fn has_schema(&self) -> bool {
!self.fields.is_empty()
}
pub fn register_dictionary(&mut self, id: i64, values: Vec<String>) {
self.encoder.register_dictionary(id, values);
}
pub fn finish(&mut self, out: &mut B) -> io::Result<()> {
out.extend_from_slice(&0xFFFF_FFFFu32.to_le_bytes());
out.extend_from_slice(&0u32.to_le_bytes());
Ok(())
}
}
impl Encoder for ArrowIpcCodec<Vec64<u8>> {
type Input = Table;
type Error = io::Error;
fn encode(&mut self, table: &Table) -> io::Result<Vec64<u8>> {
let mut out: Vec64<u8> = Vec64::new();
let view = minarrow::TableV::from_table(table.clone(), 0, table.n_rows);
self.encode_stream_batch(&view, &mut out, 0, None)?;
self.finish(&mut out)?;
Ok(out)
}
}
impl Decoder for ArrowIpcCodec<Vec64<u8>> {
type Output = Table;
type Error = io::Error;
fn decode(&mut self, bytes: &[u8]) -> io::Result<Table> {
self.decode_owned(Vec64::from_slice(bytes))
}
fn decode_owned(&mut self, bytes: Vec64<u8>) -> io::Result<Table> {
let payload = SharedBuffer::from_vec64(bytes);
self.decode_payload(payload)
}
}