Skip to main content

ahri_tre_tabular/
streaming.rs

1//! Incremental representation writers. Each call retains only the current batch;
2//! Parquet flushes each bounded batch as a row group.
3use 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
11/// Checks lengths before Arrow allocation, retains no decoded batches, and
12/// requires the canonical stream end marker (the Arrow decoder alone does not).
13pub 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                    // This transport emits uncompressed IPC. Reject compression
108                    // before a decoder can allocate from an untrusted expansion.
109                    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}
234// Enforce the encoded-message and cumulative-dictionary limits at the producer
235// too. A representation which cannot fit never emits a successful terminal.
236struct 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        // Engine metadata is not public content identity. Retain declared field
272        // names/types/nullability only; recursively strip child metadata too.
273        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}