Skip to main content

ahri_tre_tabular/
streaming_parquet.rs

1//! Standard Parquet with versioned validation metadata in padding before each
2//! row group. Ordinary readers use the final footer; incremental readers can
3//! decode each group without retaining the entire artifact or seeking backward.
4use 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        // A mini artifact has exactly one bounded group (or just its schema).
73        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
154// Restrict this producer profile before invoking a general-purpose decoder.
155// Uncompressed, non-dictionary pages avoid decompression/dictionary expansion.
156// Nested level/value counts are bounded before allocation.
157fn 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
265// Arrow schema metadata can expand null fixed-size lists without corresponding
266// physical values. Bound dimensions and nested products before reader creation.
267fn 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
323// This profile emits only scalar V1 page headers without statistics. Reading
324// just that closed subset avoids allocating from arbitrary Thrift string/list
325// lengths before the ordinary Parquet decoder receives the bounded page.
326fn 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                // PLAIN values, RLE definition and repetition levels.
365                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}