Skip to main content

ahri_tre_tabular/
ipc.rs

1use std::io::{Cursor, Write};
2
3use arrow_ipc::reader::{FileReader, StreamReader};
4use arrow_ipc::writer::{FileWriter, IpcWriteOptions, StreamWriter};
5
6use crate::{ArrowTable, Result};
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub enum ArrowIpcCompression {
10    Zstd,
11}
12
13#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
14pub struct ArrowIpcWriteOptions {
15    pub compression: Option<ArrowIpcCompression>,
16}
17
18pub fn write_ipc_stream(table: &ArrowTable) -> Result<Vec<u8>> {
19    write_ipc_stream_with_options(table, ArrowIpcWriteOptions::default())
20}
21
22pub fn write_ipc_stream_with_options(
23    table: &ArrowTable,
24    options: ArrowIpcWriteOptions,
25) -> Result<Vec<u8>> {
26    let mut buffer = Vec::new();
27    write_ipc_stream_to_writer(&mut buffer, table, options)?;
28    Ok(buffer)
29}
30
31pub fn write_ipc_stream_to_writer<W: Write>(
32    writer: W,
33    table: &ArrowTable,
34    options: ArrowIpcWriteOptions,
35) -> Result<()> {
36    let options = ipc_write_options(options)?;
37    let mut writer = StreamWriter::try_new_with_options(writer, table.schema().as_ref(), options)?;
38    for batch in table.batches() {
39        writer.write(batch)?;
40    }
41    writer.finish()?;
42    Ok(())
43}
44
45pub fn read_ipc_stream(bytes: &[u8]) -> Result<ArrowTable> {
46    let mut reader = StreamReader::try_new(Cursor::new(bytes), None)?;
47    let schema = reader.schema();
48    let mut batches = Vec::new();
49    for batch in &mut reader {
50        batches.push(batch?);
51    }
52    ArrowTable::try_new(schema, batches)
53}
54
55pub fn write_ipc_file(table: &ArrowTable) -> Result<Vec<u8>> {
56    write_ipc_file_with_options(table, ArrowIpcWriteOptions::default())
57}
58
59pub fn write_ipc_file_with_options(
60    table: &ArrowTable,
61    options: ArrowIpcWriteOptions,
62) -> Result<Vec<u8>> {
63    let mut buffer = Vec::new();
64    write_ipc_file_to_writer(&mut buffer, table, options)?;
65    Ok(buffer)
66}
67
68pub fn write_ipc_file_to_writer<W: Write>(
69    writer: W,
70    table: &ArrowTable,
71    options: ArrowIpcWriteOptions,
72) -> Result<()> {
73    let options = ipc_write_options(options)?;
74    let mut writer = FileWriter::try_new_with_options(writer, table.schema().as_ref(), options)?;
75    for batch in table.batches() {
76        writer.write(batch)?;
77    }
78    writer.finish()?;
79    Ok(())
80}
81
82pub fn read_ipc_file(bytes: &[u8]) -> Result<ArrowTable> {
83    let reader = FileReader::try_new(Cursor::new(bytes), None)?;
84    let schema = reader.schema();
85    let mut batches = Vec::new();
86    for batch in reader {
87        batches.push(batch?);
88    }
89    ArrowTable::try_new(schema, batches)
90}
91
92fn ipc_write_options(options: ArrowIpcWriteOptions) -> Result<IpcWriteOptions> {
93    let compression = options
94        .compression
95        .map(|ArrowIpcCompression::Zstd| arrow_ipc::CompressionType::ZSTD);
96    Ok(IpcWriteOptions::default().try_with_compression(compression)?)
97}
98
99/// Opens the bounded message region of an IPC file or stream. File footers are
100/// validated before using any untrusted footer length; batches remain streaming.
101pub fn open_ipc_messages(path: &std::path::Path) -> std::io::Result<std::io::Take<std::fs::File>> {
102    use std::io::{Read, Seek, SeekFrom};
103    let invalid = || std::io::Error::other("Invalid Arrow IPC container");
104    let mut file = std::fs::File::open(path)?;
105    let length = file.metadata()?.len();
106    let mut magic = [0; 8];
107    file.read_exact(&mut magic)?;
108    if &magic[..6] == b"ARROW1" {
109        if length < 18 || magic[6..] != [0, 0] {
110            return Err(invalid());
111        }
112        file.seek(SeekFrom::End(-10))?;
113        let mut tail = [0; 10];
114        file.read_exact(&mut tail)?;
115        if &tail[4..] != b"ARROW1" {
116            return Err(invalid());
117        }
118        let size = u32::from_le_bytes(tail[..4].try_into().unwrap()) as usize;
119        if size > crate::streaming::MAX_BATCH_BYTES || size as u64 + 18 > length {
120            return Err(invalid());
121        }
122        let end = length - size as u64 - 10;
123        file.seek(SeekFrom::Start(end))?;
124        let mut footer = vec![0; size];
125        file.read_exact(&mut footer)?;
126        arrow_ipc::root_as_footer(&footer).map_err(|_| invalid())?;
127        // FileWriter pads the six-byte magic to its configured 8/16/32/64
128        // byte alignment. Padding is not an empty stream terminator.
129        let mut start = 8;
130        loop {
131            if start >= end || start > 64 {
132                return Err(invalid());
133            }
134            file.seek(SeekFrom::Start(start))?;
135            let mut prefix = [0; 8];
136            file.read_exact(&mut prefix)?;
137            if prefix != [0; 8] {
138                break;
139            }
140            start += 8;
141        }
142        file.seek(SeekFrom::Start(start))?;
143        Ok(file.take(end - start))
144    } else {
145        file.seek(SeekFrom::Start(0))?;
146        Ok(file.take(length))
147    }
148}
149
150pub fn read_ipc_schema(path: &std::path::Path) -> Result<crate::SchemaRef> {
151    Ok(StreamReader::try_new(open_ipc_messages(path)?, None)?.schema())
152}