1use super::{MAX_BATCH_BYTES, invalid};
5use crate::{RecordBatch, SchemaRef};
6use bytes::Bytes;
7use parquet::{
8 arrow::{ArrowSchemaConverter, ArrowWriter, arrow_reader::ParquetRecordBatchReaderBuilder},
9 basic::{Compression, Encoding},
10 column::writer::ColumnCloseResult,
11 file::{
12 metadata::{ParquetMetaData, ParquetMetaDataReader, RowGroupMetaData},
13 properties::{EnabledStatistics, WriterProperties},
14 writer::SerializedFileWriter,
15 },
16};
17use sha2::{Digest, Sha256};
18use std::{
19 io::{self, Write},
20 sync::Arc,
21};
22
23const GROUP: &[u8; 8] = b"ATPQ0001";
24const END: &[u8; 8] = b"ATPQEND1";
25#[path = "parquet_metadata_bounds.rs"]
26mod metadata_bounds;
27
28fn properties() -> WriterProperties {
29 WriterProperties::builder()
30 .set_dictionary_enabled(false)
31 .set_statistics_enabled(EnabledStatistics::None)
32 .set_offset_index_disabled(true)
33 .set_max_row_group_row_count(Some(8192))
34 .set_data_page_size_limit(1024 * 1024)
35 .build()
36}
37struct Buffer(Vec<u8>);
38impl Write for Buffer {
39 fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
40 if self.0.len().saturating_add(bytes.len()) > MAX_BATCH_BYTES {
41 return Err(invalid());
42 }
43 self.0.extend_from_slice(bytes);
44 Ok(bytes.len())
45 }
46 fn flush(&mut self) -> io::Result<()> {
47 Ok(())
48 }
49}
50pub(super) struct Writer<W: Write + Send> {
51 writer: SerializedFileWriter<W>,
52 schema: SchemaRef,
53 metadata_bytes: usize,
54}
55impl<W: Write + Send> Writer<W> {
56 pub(super) fn new(sink: W, schema: SchemaRef) -> io::Result<Self> {
57 let parquet_schema = ArrowSchemaConverter::new()
58 .convert(&schema)
59 .map_err(|_| invalid())?;
60 let mut props = properties();
61 parquet::arrow::add_encoded_arrow_schema_to_metadata(&schema, &mut props);
62 let writer =
63 SerializedFileWriter::new(sink, parquet_schema.root_schema_ptr(), Arc::new(props))
64 .map_err(|_| invalid())?;
65 Ok(Self {
66 writer,
67 schema,
68 metadata_bytes: 0,
69 })
70 }
71 pub(super) fn write(&mut self, batch: &RecordBatch) -> io::Result<()> {
72 if batch.num_rows() > 8192 {
74 return Err(invalid());
75 }
76 let mut mini =
77 ArrowWriter::try_new(Buffer(Vec::new()), self.schema.clone(), Some(properties()))
78 .map_err(|_| invalid())?;
79 mini.write(batch).map_err(|_| invalid())?;
80 let bytes = Bytes::from(mini.into_inner().map_err(|_| invalid())?.0);
81 let (footer, metadata) = footer(&bytes)?;
82 let body = bytes.len() - footer.len() - 12;
83 self.metadata_bytes = self
84 .metadata_bytes
85 .checked_add(footer.len())
86 .filter(|n| *n <= MAX_BATCH_BYTES)
87 .ok_or_else(invalid)?;
88 validate_group(&bytes, &metadata, u64::MAX)?;
89 self.writer.write_all(GROUP)?;
90 self.writer
91 .write_all(&(footer.len() as u32).to_le_bytes())?;
92 self.writer.write_all(&(body as u32).to_le_bytes())?;
93 self.writer.write_all(footer)?;
94 for group in metadata.row_groups() {
95 let mut target = self.writer.next_row_group().map_err(|_| invalid())?;
96 for column in group.columns() {
97 target
98 .append_column(
99 &bytes,
100 ColumnCloseResult {
101 bytes_written: column.compressed_size() as u64,
102 rows_written: group.num_rows() as u64,
103 metadata: column.clone(),
104 bloom_filter: None,
105 column_index: None,
106 offset_index: None,
107 },
108 )
109 .map_err(|_| invalid())?;
110 }
111 target.close().map_err(|_| invalid())?;
112 }
113 Ok(())
114 }
115 pub(super) fn finish(mut self) -> io::Result<W> {
116 if self.metadata_bytes == 0 {
117 self.write(&RecordBatch::new_empty(self.schema.clone()))?;
118 }
119 self.writer.write_all(END)?;
120 self.writer.into_inner().map_err(|_| invalid())
121 }
122}
123
124fn footer_bytes(bytes: &[u8]) -> io::Result<&[u8]> {
125 if bytes.len() < 8 || !bytes.ends_with(b"PAR1") {
126 return Err(invalid());
127 }
128 let length =
129 u32::from_le_bytes(bytes[bytes.len() - 8..bytes.len() - 4].try_into().unwrap()) as usize;
130 if length > MAX_BATCH_BYTES || length > bytes.len() - 8 {
131 return Err(invalid());
132 }
133 Ok(&bytes[bytes.len() - 8 - length..bytes.len() - 8])
134}
135fn footer(bytes: &[u8]) -> io::Result<(&[u8], Arc<ParquetMetaData>)> {
136 let footer = footer_bytes(bytes)?;
137 Ok((footer, decode_metadata(footer)?))
138}
139fn decode_metadata(bytes: &[u8]) -> io::Result<Arc<ParquetMetaData>> {
140 metadata_bounds::check(bytes)?;
141 ParquetMetaDataReader::decode_metadata(bytes)
142 .map(Arc::new)
143 .map_err(|_| invalid())
144}
145fn arrow_schema(metadata: &ParquetMetaData) -> io::Result<&str> {
146 metadata
147 .file_metadata()
148 .key_value_metadata()
149 .and_then(|values| values.iter().find(|value| value.key == "ARROW:schema"))
150 .and_then(|value| value.value.as_deref())
151 .ok_or_else(invalid)
152}
153
154fn validate_group(
158 bytes: &Bytes,
159 metadata: &Arc<ParquetMetaData>,
160 max_rows: u64,
161) -> io::Result<u64> {
162 let schema = metadata.file_metadata().schema_descr();
163 let rows = u64::try_from(metadata.file_metadata().num_rows()).map_err(|_| invalid())?;
164 if rows > max_rows.min(8192) || schema.num_columns() > 512 || metadata.num_row_groups() > 1 {
165 return Err(invalid());
166 }
167 let footer = footer_bytes(bytes)?;
168 let end = bytes.len() - footer.len() - 8;
169 if end.saturating_add(schema.num_columns() * 4096) > MAX_BATCH_BYTES {
170 return Err(invalid());
171 }
172 let mut position = 4usize;
173 let mut decoded_bound = end;
174 for group in metadata.row_groups() {
175 if group.num_rows() as u64 != rows {
176 return Err(invalid());
177 }
178 for column in group.columns() {
179 let desc = column.column_descr();
180 let (offset, length) = column.byte_range();
181 if desc.max_rep_level() > 16
182 || desc.max_def_level() > 32
183 || desc.type_length() > MAX_BATCH_BYTES as i32
184 || column.compression() != Compression::UNCOMPRESSED
185 || column.file_path().is_some()
186 || column.dictionary_page_offset().is_some()
187 || column.num_values() < 0
188 || column
189 .encodings()
190 .any(|value| !matches!(value, Encoding::PLAIN | Encoding::RLE))
191 || offset != position as u64
192 || length > MAX_BATCH_BYTES as u64
193 || offset.checked_add(length).is_none_or(|n| n > end as u64)
194 {
195 return Err(invalid());
196 }
197 let width = desc.type_length().max(16) as usize
198 + (desc.max_def_level() + desc.max_rep_level()) as usize * 8;
199 decoded_bound = usize::try_from(column.num_values())
200 .ok()
201 .and_then(|n| n.checked_mul(width))
202 .and_then(|n| decoded_bound.checked_add(n))
203 .filter(|n| *n <= MAX_BATCH_BYTES)
204 .ok_or_else(invalid)?;
205 let column_end = position + length as usize;
206 let mut values = 0u64;
207 while position < column_end {
208 let mut cursor = std::io::Cursor::new(&bytes[position..column_end]);
209 let (page_size, page_values) = page_header(&mut cursor)?;
210 let header_size = cursor.position() as usize;
211 if header_size > 16384 {
212 return Err(invalid());
213 }
214 values = values
215 .checked_add(page_values)
216 .filter(|v| *v <= column.num_values() as u64)
217 .ok_or_else(invalid)?;
218 position = position
219 .checked_add(header_size)
220 .and_then(|n| n.checked_add(page_size))
221 .filter(|n| *n <= column_end)
222 .ok_or_else(invalid)?;
223 }
224 if values != column.num_values() as u64 {
225 return Err(invalid());
226 }
227 }
228 }
229 if position != end {
230 return Err(invalid());
231 }
232 let metadata = parquet::arrow::arrow_reader::ArrowReaderMetadata::try_new(
233 metadata.clone(),
234 Default::default(),
235 )
236 .map_err(|_| invalid())?;
237 let entries = metadata
238 .metadata()
239 .row_groups()
240 .iter()
241 .flat_map(|group| group.columns())
242 .map(|column| column.num_values() as u64)
243 .max()
244 .unwrap_or(rows)
245 .max(rows);
246 bound_arrow_expansion(metadata.schema(), entries, end)?;
247 let reader = ParquetRecordBatchReaderBuilder::new_with_metadata(bytes.clone(), metadata)
248 .with_batch_size(32)
249 .build()
250 .map_err(|_| invalid())?;
251 let mut decoded = 0u64;
252 for batch in reader {
253 let batch = batch.map_err(|_| invalid())?;
254 if batch.get_array_memory_size() > MAX_BATCH_BYTES {
255 return Err(invalid());
256 }
257 decoded += batch.num_rows() as u64;
258 }
259 if decoded != rows {
260 return Err(invalid());
261 }
262 Ok(rows)
263}
264
265fn bound_arrow_expansion(schema: &crate::Schema, entries: u64, encoded: usize) -> io::Result<()> {
268 use crate::DataType;
269 fn width(data: &DataType, depth: usize) -> io::Result<usize> {
270 if depth > 32 {
271 return Err(invalid());
272 }
273 let add = |a: usize, b: usize| {
274 a.checked_add(b)
275 .filter(|n| *n <= MAX_BATCH_BYTES)
276 .ok_or_else(invalid)
277 };
278 match data {
279 DataType::FixedSizeList(field, length) => usize::try_from(*length)
280 .ok()
281 .and_then(|length| length.checked_mul(width(field.data_type(), depth + 1).ok()?))
282 .and_then(|n| n.checked_add(16))
283 .filter(|n| *n <= MAX_BATCH_BYTES)
284 .ok_or_else(invalid),
285 DataType::FixedSizeBinary(length) => usize::try_from(*length)
286 .ok()
287 .filter(|n| *n <= MAX_BATCH_BYTES)
288 .ok_or_else(invalid),
289 DataType::Struct(fields) => fields.iter().try_fold(16, |total, field| {
290 add(total, width(field.data_type(), depth + 1)?)
291 }),
292 DataType::List(field)
293 | DataType::LargeList(field)
294 | DataType::ListView(field)
295 | DataType::LargeListView(field)
296 | DataType::Map(field, _) => add(16, width(field.data_type(), depth + 1)?),
297 DataType::Dictionary(_, value) => add(16, width(value, depth + 1)?),
298 DataType::RunEndEncoded(ends, values) => add(
299 width(ends.data_type(), depth + 1)?,
300 width(values.data_type(), depth + 1)?,
301 ),
302 DataType::Union(fields, _) => fields.iter().try_fold(16, |total, (_, field)| {
303 add(total, width(field.data_type(), depth + 1)?)
304 }),
305 DataType::Decimal256(_, _) => Ok(32),
306 _ => Ok(16),
307 }
308 }
309 let row = schema.fields().iter().try_fold(0usize, |total, field| {
310 total
311 .checked_add(width(field.data_type(), 0)?)
312 .ok_or_else(invalid)
313 })?;
314 usize::try_from(entries)
315 .ok()
316 .and_then(|n| n.checked_mul(row))
317 .and_then(|n| n.checked_add(encoded))
318 .filter(|n| *n <= MAX_BATCH_BYTES)
319 .ok_or_else(invalid)?;
320 Ok(())
321}
322
323fn page_header(input: &mut impl std::io::Read) -> io::Result<(usize, u64)> {
327 use thrift::protocol::{TCompactInputProtocol, TInputProtocol, TType};
328 let error = |_| invalid();
329 let mut protocol = TCompactInputProtocol::new(input);
330 protocol.read_struct_begin().map_err(error)?;
331 let mut fields = [None; 4];
332 let mut values = None;
333 loop {
334 let field = protocol.read_field_begin().map_err(error)?;
335 if field.field_type == TType::Stop {
336 break;
337 }
338 match (field.id, field.field_type) {
339 (Some(id @ 1..=4), TType::I32) => {
340 let slot = &mut fields[(id - 1) as usize];
341 if slot.is_some() {
342 return Err(invalid());
343 }
344 *slot = Some(protocol.read_i32().map_err(error)?);
345 }
346 (Some(5), TType::Struct) if values.is_none() => {
347 protocol.read_struct_begin().map_err(error)?;
348 let mut data = [None; 4];
349 loop {
350 let field = protocol.read_field_begin().map_err(error)?;
351 if field.field_type == TType::Stop {
352 break;
353 }
354 let Some(id @ 1..=4) = field.id else {
355 return Err(invalid());
356 };
357 if field.field_type != TType::I32 || data[(id - 1) as usize].is_some() {
358 return Err(invalid());
359 }
360 data[(id - 1) as usize] = Some(protocol.read_i32().map_err(error)?);
361 protocol.read_field_end().map_err(error)?;
362 }
363 protocol.read_struct_end().map_err(error)?;
364 if data[1..] != [Some(0), Some(3), Some(3)] {
366 return Err(invalid());
367 }
368 values = Some(u64::try_from(data[0].ok_or_else(invalid)?).map_err(|_| invalid())?);
369 }
370 _ => return Err(invalid()),
371 }
372 protocol.read_field_end().map_err(error)?;
373 }
374 protocol.read_struct_end().map_err(error)?;
375 if fields[0] != Some(0) || fields[1] != fields[2] {
376 return Err(invalid());
377 }
378 let length = usize::try_from(fields[2].ok_or_else(invalid)?).map_err(|_| invalid())?;
379 if length > MAX_BATCH_BYTES {
380 return Err(invalid());
381 }
382 Ok((length, values.ok_or_else(invalid)?))
383}
384
385fn hash_group(hash: &mut Sha256, group: &RowGroupMetaData, base: u64) -> io::Result<()> {
386 hash.update(group.num_rows().to_le_bytes());
387 hash.update((group.num_columns() as u64).to_le_bytes());
388 for column in group.columns() {
389 let (offset, size) = column.byte_range();
390 hash.update(offset.checked_add(base).ok_or_else(invalid)?.to_le_bytes());
391 hash.update(size.to_le_bytes());
392 hash.update(column.num_values().to_le_bytes());
393 hash.update(column.uncompressed_size().to_le_bytes());
394 hash.update(format!("{:?}:{:?}", column.column_descr(), column.compression()).as_bytes());
395 for encoding in column.encodings() {
396 hash.update([encoding as u8]);
397 }
398 if column.file_path().is_some() || column.dictionary_page_offset().is_some() {
399 return Err(invalid());
400 }
401 }
402 Ok(())
403}
404
405pub struct ParquetValidator {
406 pending: Vec<u8>,
407 target: usize,
408 position: u64,
409 started: bool,
410 group: Option<(usize, usize)>,
411 lengths: bool,
412 ending: bool,
413 schema: Option<Arc<parquet::schema::types::SchemaDescriptor>>,
414 arrow_schema: Option<String>,
415 metadata_bytes: usize,
416 rows: u64,
417 max_rows: u64,
418 groups: usize,
419 hash: Sha256,
420}
421impl ParquetValidator {
422 pub fn new(max_rows: u64) -> Self {
423 Self {
424 pending: Vec::new(),
425 target: 4,
426 position: 0,
427 started: false,
428 group: None,
429 lengths: false,
430 ending: false,
431 schema: None,
432 arrow_schema: None,
433 metadata_bytes: 0,
434 rows: 0,
435 max_rows,
436 groups: 0,
437 hash: Sha256::new(),
438 }
439 }
440 pub fn push(&mut self, mut bytes: &[u8]) -> io::Result<()> {
441 while !bytes.is_empty() {
442 if self.ending {
443 if self.pending.len().saturating_add(bytes.len()) > MAX_BATCH_BYTES {
444 return Err(invalid());
445 }
446 self.pending.extend_from_slice(bytes);
447 return Ok(());
448 }
449 let size = bytes.len().min(self.target - self.pending.len());
450 self.pending.extend_from_slice(&bytes[..size]);
451 bytes = &bytes[size..];
452 if self.pending.len() != self.target {
453 continue;
454 }
455 self.position = self
456 .position
457 .checked_add(self.pending.len() as u64)
458 .ok_or_else(invalid)?;
459 if !self.started {
460 if self.pending != b"PAR1" {
461 return Err(invalid());
462 }
463 self.started = true;
464 self.pending.clear();
465 self.target = 8;
466 continue;
467 }
468 if let Some((footer_len, body_len)) = self.group.take() {
469 let metadata = decode_metadata(&self.pending[..footer_len])?;
470 let mut mini = Vec::with_capacity(body_len + footer_len + 12);
471 mini.extend_from_slice(b"PAR1");
472 mini.extend_from_slice(&self.pending[footer_len..]);
473 mini.extend_from_slice(&self.pending[..footer_len]);
474 mini.extend_from_slice(&(footer_len as u32).to_le_bytes());
475 mini.extend_from_slice(b"PAR1");
476 let mini = Bytes::from(mini);
477 let rows =
478 validate_group(&mini, &metadata, self.max_rows.saturating_sub(self.rows))?;
479 if self.schema.as_ref().is_some_and(|schema| {
480 schema.as_ref() != metadata.file_metadata().schema_descr()
481 }) {
482 return Err(invalid());
483 }
484 self.schema = Some(metadata.file_metadata().schema_descr_ptr());
485 let arrow_schema = arrow_schema(&metadata)?;
486 if self
487 .arrow_schema
488 .as_ref()
489 .is_some_and(|schema| schema != arrow_schema)
490 {
491 return Err(invalid());
492 }
493 self.arrow_schema = Some(arrow_schema.to_owned());
494 for group in metadata.row_groups() {
495 hash_group(&mut self.hash, group, self.position - body_len as u64 - 4)?;
496 self.groups += 1;
497 }
498 self.rows = self.rows.checked_add(rows).ok_or_else(invalid)?;
499 self.pending.clear();
500 self.target = 8;
501 } else if !self.lengths {
502 if self.pending == END {
503 self.ending = true;
504 self.pending.clear();
505 } else if self.pending == GROUP {
506 self.pending.clear();
507 self.lengths = true;
508 } else {
509 return Err(invalid());
510 }
511 } else {
512 let footer_len = u32::from_le_bytes(self.pending[..4].try_into().unwrap()) as usize;
513 let body_len = u32::from_le_bytes(self.pending[4..8].try_into().unwrap()) as usize;
514 if footer_len < 8
515 || footer_len.saturating_add(body_len).saturating_add(12) > MAX_BATCH_BYTES
516 {
517 return Err(invalid());
518 }
519 self.metadata_bytes = self
520 .metadata_bytes
521 .checked_add(footer_len)
522 .filter(|n| *n <= MAX_BATCH_BYTES)
523 .ok_or_else(invalid)?;
524 self.pending.clear();
525 self.lengths = false;
526 self.group = Some((footer_len, body_len));
527 self.target = footer_len + body_len;
528 }
529 }
530 Ok(())
531 }
532 pub fn finish(&mut self) -> io::Result<u64> {
533 if !self.ending {
534 return Err(invalid());
535 }
536 let (footer, metadata) = footer(&self.pending)?;
537 if footer.len() + 8 != self.pending.len() {
538 return Err(invalid());
539 }
540 if self.arrow_schema.as_deref() != Some(arrow_schema(&metadata)?) {
541 return Err(invalid());
542 }
543 if self
544 .schema
545 .as_ref()
546 .is_none_or(|schema| schema.as_ref() != metadata.file_metadata().schema_descr())
547 || metadata.file_metadata().num_rows() < 0
548 || metadata.file_metadata().num_rows() as u64 != self.rows
549 || metadata.num_row_groups() != self.groups
550 {
551 return Err(invalid());
552 }
553 let mut hash = Sha256::new();
554 for group in metadata.row_groups() {
555 hash_group(&mut hash, group, 0)?;
556 }
557 if hash.finalize() != self.hash.clone().finalize() {
558 return Err(invalid());
559 }
560 Ok(self.rows)
561 }
562}
563
564#[cfg(test)]
565mod tests {
566 #[test]
567 fn null_fixed_size_list_dimensions_are_bounded_before_decode() {
568 use crate::{DataType, Field, Schema};
569 use std::sync::Arc;
570 let item = Arc::new(Field::new("item", DataType::Int64, true));
571 let schema = Schema::new(vec![Field::new(
572 "null_list",
573 DataType::FixedSizeList(item.clone(), 1_000_000_000),
574 true,
575 )]);
576 assert!(super::bound_arrow_expansion(&schema, 1, 32).is_err());
577 let child = Arc::new(Field::new(
578 "child",
579 DataType::FixedSizeList(item.clone(), 1000),
580 true,
581 ));
582 let nested = Schema::new(vec![Field::new(
583 "nested",
584 DataType::FixedSizeList(child, 1000),
585 true,
586 )]);
587 assert!(super::bound_arrow_expansion(&nested, 1, 32).is_err());
588 let small = Schema::new(vec![Field::new(
589 "bounded",
590 DataType::FixedSizeList(item, 4),
591 true,
592 )]);
593 assert!(super::bound_arrow_expansion(&small, 32, 1024).is_ok());
594 }
595}