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
99pub 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 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}