1use crate::{RecordBatch, Schema, SchemaRef};
4use std::{
5 io::{self, Write},
6 sync::Arc,
7};
8
9pub const MAX_BATCH_BYTES: usize = 8 * 1024 * 1024;
10
11pub struct ArrowValidator {
14 allow_compression: bool,
15 decoded_bytes: u64,
16 max_decoded_bytes: u64,
17 decoder: arrow_ipc::reader::StreamDecoder,
18 message: Vec<u8>,
19 target: usize,
20 metadata: Option<usize>,
21 body_checked: bool,
22 dictionaries: usize,
23 rows: u64,
24 max_rows: u64,
25 ended: bool,
26}
27impl ArrowValidator {
28 pub fn new(max_rows: u64) -> Self {
29 Self {
30 allow_compression: false,
31 decoded_bytes: 0,
32 max_decoded_bytes: u64::MAX,
33 decoder: arrow_ipc::reader::StreamDecoder::new(),
34 message: Vec::new(),
35 target: 8,
36 metadata: None,
37 body_checked: false,
38 dictionaries: 0,
39 rows: 0,
40 max_rows,
41 ended: false,
42 }
43 }
44 pub fn for_upload(max_rows: u64, max_decoded_bytes: u64) -> Self {
45 Self {
46 allow_compression: true,
47 max_decoded_bytes,
48 ..Self::new(max_rows)
49 }
50 }
51 pub fn push(&mut self, mut bytes: &[u8]) -> io::Result<()> {
52 while !bytes.is_empty() {
53 if self.ended {
54 return Err(invalid());
55 }
56 let size = bytes.len().min(self.target - self.message.len());
57 self.message.extend_from_slice(&bytes[..size]);
58 bytes = &bytes[size..];
59 if self.message.len() != self.target {
60 continue;
61 }
62 if self.metadata.is_none() {
63 if self.message[..4] != [255; 4] {
64 return Err(invalid());
65 }
66 let metadata = u32::from_le_bytes(self.message[4..8].try_into().unwrap()) as usize;
67 if metadata > MAX_BATCH_BYTES - 8 {
68 return Err(invalid());
69 }
70 if metadata == 0 {
71 self.decode_message()?;
72 if self.decoder.schema().is_none() {
73 return Err(invalid());
74 }
75 self.ended = true;
76 continue;
77 }
78 self.target = 8 + metadata;
79 self.metadata = Some(metadata);
80 continue;
81 }
82 let metadata = self.metadata.unwrap();
83 if !self.body_checked {
84 let message = arrow_ipc::root_as_message(&self.message[8..8 + metadata])
85 .map_err(|_| invalid())?;
86 let body: usize = message.bodyLength().try_into().map_err(|_| invalid())?;
87 self.target = (8 + metadata)
88 .checked_add(body)
89 .filter(|size| *size <= MAX_BATCH_BYTES)
90 .ok_or_else(invalid)?;
91 let batch = match message.header_type() {
92 arrow_ipc::MessageHeader::Schema => None,
93 arrow_ipc::MessageHeader::RecordBatch => message.header_as_record_batch(),
94 arrow_ipc::MessageHeader::DictionaryBatch => {
95 self.dictionaries = self
96 .dictionaries
97 .checked_add(self.target)
98 .filter(|size| *size <= MAX_BATCH_BYTES)
99 .ok_or_else(invalid)?;
100 message
101 .header_as_dictionary_batch()
102 .and_then(|batch| batch.data())
103 }
104 _ => return Err(invalid()),
105 };
106 if let Some(batch) = batch {
107 if (batch.compression().is_some() && !self.allow_compression)
110 || batch.length() < 0
111 || batch.length() as u64 > self.max_rows
112 {
113 return Err(invalid());
114 }
115 if let Some(nodes) = batch.nodes()
116 && nodes
117 .iter()
118 .any(|node| node.length() < 0 || node.length() as u64 > self.max_rows)
119 {
120 return Err(invalid());
121 }
122 }
123 self.body_checked = true;
124 if body != 0 {
125 continue;
126 }
127 }
128 self.decode_message()?;
129 self.target = 8;
130 self.metadata = None;
131 self.body_checked = false;
132 }
133 Ok(())
134 }
135 fn decode_message(&mut self) -> io::Result<()> {
136 let mut decoded_size = self.message.len();
137 if self.allow_compression
138 && let Some(metadata) = self.metadata
139 {
140 let message = arrow_ipc::root_as_message(&self.message[8..8 + metadata])
141 .map_err(|_| invalid())?;
142 let batch = message
143 .header_as_record_batch()
144 .or_else(|| message.header_as_dictionary_batch().and_then(|b| b.data()));
145 if let Some(batch) = batch
146 && batch.compression().is_some()
147 {
148 let body = &self.message[8 + metadata..];
149 let mut decoded = 0usize;
150 for buffer in batch.buffers().ok_or_else(invalid)?.iter() {
151 let start: usize = buffer.offset().try_into().map_err(|_| invalid())?;
152 let size: usize = buffer.length().try_into().map_err(|_| invalid())?;
153 if size == 0 {
154 continue;
155 }
156 let data = body
157 .get(start..start.checked_add(size).ok_or_else(invalid)?)
158 .ok_or_else(invalid)?;
159 let length =
160 i64::from_le_bytes(data.get(..8).ok_or_else(invalid)?.try_into().unwrap());
161 let length = if length == -1 {
162 size - 8
163 } else {
164 usize::try_from(length).map_err(|_| invalid())?
165 };
166 decoded = decoded
167 .checked_add(length)
168 .filter(|n| *n <= MAX_BATCH_BYTES)
169 .ok_or_else(invalid)?;
170 }
171 decoded_size = decoded_size.max(8 + metadata + decoded);
172 if message.header_as_dictionary_batch().is_some() {
173 self.dictionaries = self
174 .dictionaries
175 .checked_add(decoded)
176 .filter(|n| *n <= MAX_BATCH_BYTES)
177 .ok_or_else(invalid)?;
178 }
179 }
180 }
181 self.decoded_bytes = self
182 .decoded_bytes
183 .checked_add(decoded_size as u64)
184 .filter(|n| *n <= self.max_decoded_bytes)
185 .ok_or_else(invalid)?;
186 let mut buffer = arrow_buffer::Buffer::from(std::mem::take(&mut self.message));
187 while !buffer.is_empty() {
188 if let Some(batch) = self.decoder.decode(&mut buffer).map_err(|_| invalid())? {
189 self.rows = self
190 .rows
191 .checked_add(batch.num_rows() as u64)
192 .filter(|rows| *rows <= self.max_rows)
193 .ok_or_else(invalid)?;
194 if batch.get_array_memory_size() > MAX_BATCH_BYTES {
195 return Err(invalid());
196 }
197 }
198 }
199 if self
200 .decoder
201 .schema()
202 .is_some_and(|schema| schema.fields().len() > 512)
203 {
204 return Err(invalid());
205 }
206 Ok(())
207 }
208 pub fn finish(&mut self) -> io::Result<u64> {
209 if !self.ended || !self.message.is_empty() {
210 return Err(invalid());
211 }
212 self.decoder.finish().map_err(|_| invalid())?;
213 Ok(self.rows)
214 }
215}
216
217pub enum StreamFormat {
218 Arrow,
219 Parquet,
220 Csv,
221 Json,
222 PreviewJson { limit: u64 },
223}
224
225#[path = "streaming_parquet.rs"]
226mod streaming_parquet;
227pub use streaming_parquet::ParquetValidator;
228enum Encoder<W: Write + Send> {
229 Arrow(arrow_ipc::writer::StreamWriter<ValidatedArrow<W>>),
230 Parquet(streaming_parquet::Writer<W>),
231 Csv(arrow_csv::Writer<W>),
232 Json(arrow_json::ArrayWriter<W>),
233}
234struct ValidatedArrow<W> {
237 sink: W,
238 validator: ArrowValidator,
239}
240impl<W: Write> Write for ValidatedArrow<W> {
241 fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
242 self.validator.push(bytes)?;
243 self.sink.write_all(bytes)?;
244 Ok(bytes.len())
245 }
246 fn flush(&mut self) -> io::Result<()> {
247 self.sink.flush()
248 }
249}
250pub struct BatchWriter<W: Write + Send> {
251 encoder: Encoder<W>,
252 schema: SchemaRef,
253 rows: u64,
254 max_rows: u64,
255 input_rows: u64,
256 preview_limit: Option<u64>,
257}
258fn invalid() -> io::Error {
259 io::Error::other("Tabular representation could not be completed")
260}
261impl<W: Write + Send> BatchWriter<W> {
262 pub fn new(
263 mut sink: W,
264 schema: &Schema,
265 format: StreamFormat,
266 max_rows: u64,
267 ) -> io::Result<Self> {
268 if schema.fields().len() > 512 {
269 return Err(invalid());
270 }
271 let value = serde_json::to_value(schema).map_err(|_| invalid())?;
274 fn strip(mut value: serde_json::Value) -> serde_json::Value {
275 match &mut value {
276 serde_json::Value::Object(map) => {
277 if map.contains_key("metadata") {
278 map.insert("metadata".into(), serde_json::json!({}));
279 }
280 for item in map.values_mut() {
281 *item = strip(item.take());
282 }
283 }
284 serde_json::Value::Array(items) => {
285 for item in items {
286 *item = strip(item.take());
287 }
288 }
289 _ => (),
290 }
291 value
292 }
293 let schema: Schema = serde_json::from_value(strip(value)).map_err(|_| invalid())?;
294 let schema = Arc::new(schema);
295 let preview_limit = match format {
296 StreamFormat::PreviewJson { limit } => Some(limit),
297 _ => None,
298 };
299 let encoder = match format {
300 StreamFormat::Arrow => Encoder::Arrow(
301 arrow_ipc::writer::StreamWriter::try_new(
302 ValidatedArrow {
303 sink,
304 validator: ArrowValidator::new(max_rows),
305 },
306 &schema,
307 )
308 .map_err(|_| invalid())?,
309 ),
310 StreamFormat::Parquet => {
311 Encoder::Parquet(streaming_parquet::Writer::new(sink, schema.clone())?)
312 }
313 StreamFormat::Csv => Encoder::Csv(arrow_csv::Writer::new(sink)),
314 StreamFormat::Json => Encoder::Json(arrow_json::ArrayWriter::new(sink)),
315 StreamFormat::PreviewJson { .. } => {
316 sink.write_all(b"{\"columns\":")?;
317 let columns = schema.fields().iter().map(|field| {
318 serde_json::json!({"name": field.name(), "value_type": field.data_type().to_string()})
319 }).collect::<Vec<_>>();
320 serde_json::to_writer(&mut sink, &columns).map_err(|_| invalid())?;
321 sink.write_all(b",\"rows\":")?;
322 Encoder::Json(arrow_json::ArrayWriter::new(sink))
323 }
324 };
325 Ok(Self {
326 encoder,
327 schema,
328 rows: 0,
329 max_rows,
330 input_rows: 0,
331 preview_limit,
332 })
333 }
334 pub fn write(&mut self, batch: &RecordBatch) -> io::Result<()> {
335 self.input_rows = self
336 .input_rows
337 .checked_add(batch.num_rows() as u64)
338 .ok_or_else(invalid)?;
339 if batch.get_array_memory_size() > MAX_BATCH_BYTES {
340 return Err(invalid());
341 }
342 let selected = self.preview_limit.map(|limit| {
343 batch.slice(
344 0,
345 (batch.num_rows() as u64).min(limit.saturating_sub(self.rows)) as usize,
346 )
347 });
348 let batch = selected.as_ref().unwrap_or(batch);
349 let rows = self
350 .rows
351 .checked_add(batch.num_rows() as u64)
352 .ok_or_else(invalid)?;
353 if rows > self.max_rows || batch.get_array_memory_size() > MAX_BATCH_BYTES {
354 return Err(invalid());
355 }
356 let clean = RecordBatch::try_new_with_options(
357 self.schema.clone(),
358 batch.columns().to_vec(),
359 &arrow_array::RecordBatchOptions::new().with_row_count(Some(batch.num_rows())),
360 )
361 .map_err(|_| invalid())?;
362 match &mut self.encoder {
363 Encoder::Arrow(writer) => writer.write(&clean).map_err(|_| invalid())?,
364 Encoder::Parquet(writer) => {
365 writer.write(&clean)?;
366 }
367 Encoder::Csv(writer) => writer.write(&clean).map_err(|_| invalid())?,
368 Encoder::Json(writer) => writer.write(&clean).map_err(|_| invalid())?,
369 }
370 self.rows = rows;
371 Ok(())
372 }
373 pub fn rows(&self) -> u64 {
374 self.rows
375 }
376 pub fn input_rows(&self) -> u64 {
377 self.input_rows
378 }
379 pub fn finish(mut self) -> io::Result<W> {
380 match self.encoder {
381 Encoder::Arrow(mut writer) => {
382 writer.finish().map_err(|_| invalid())?;
383 let mut validated = writer.into_inner().map_err(|_| invalid())?;
384 validated.validator.finish()?;
385 Ok(validated.sink)
386 }
387 Encoder::Parquet(writer) => writer.finish(),
388 Encoder::Csv(writer) => Ok(writer.into_inner()),
389 Encoder::Json(ref mut writer) => {
390 writer.finish().map_err(|_| invalid())?;
391 if let Encoder::Json(writer) = self.encoder {
392 let mut sink = writer.into_inner();
393 if let Some(limit) = self.preview_limit {
394 write!(
395 sink,
396 ",\"displayed_row_count\":{},\"total_row_count\":{},\"truncated\":{},\"limit\":{limit}}}",
397 self.rows,
398 self.input_rows,
399 self.rows < self.input_rows
400 )?;
401 }
402 Ok(sink)
403 } else {
404 unreachable!()
405 }
406 }
407 }
408 }
409}