Skip to main content

arrow_csv/reader/
records.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18use arrow_schema::ArrowError;
19use csv_core::{ReadRecordResult, Reader};
20
21/// The estimated length of a field in bytes
22const AVERAGE_FIELD_SIZE: usize = 8;
23
24/// The minimum amount of data in a single read
25const MIN_CAPACITY: usize = 1024;
26
27/// [`RecordDecoder`] provides a push-based interface to decoder [`StringRecords`]
28#[derive(Debug)]
29pub struct RecordDecoder {
30    delimiter: Reader,
31
32    /// The expected number of fields per row
33    num_columns: usize,
34
35    /// The current line number
36    line_number: usize,
37
38    /// Offsets delimiting field start positions
39    offsets: Vec<usize>,
40
41    /// The current offset into `self.offsets`
42    ///
43    /// We track this independently of Vec to avoid re-zeroing memory
44    offsets_len: usize,
45
46    /// The number of fields read for the current record
47    current_field: usize,
48
49    /// The number of rows buffered
50    num_rows: usize,
51
52    /// Decoded field data
53    data: Vec<u8>,
54
55    /// Offsets into data
56    ///
57    /// We track this independently of Vec to avoid re-zeroing memory
58    data_len: usize,
59
60    /// Whether rows with less than expected columns are considered valid
61    ///
62    /// Default value is false
63    /// When enabled fills in missing columns with null
64    truncated_rows: bool,
65
66    /// The number of rows padded because they had fewer fields than expected
67    ///
68    /// Only incremented when `truncated_rows` is enabled. Cumulative over the
69    /// lifetime of this decoder, i.e. not reset by [`Self::flush`], but reset by
70    /// [`Self::clear`] along with the buffered rows it counted
71    truncated_row_count: usize,
72}
73
74impl RecordDecoder {
75    pub fn new(delimiter: Reader, num_columns: usize, truncated_rows: bool) -> Self {
76        Self {
77            delimiter,
78            num_columns,
79            line_number: 1,
80            offsets: vec![],
81            offsets_len: 1, // The first offset is always 0
82            current_field: 0,
83            data_len: 0,
84            data: vec![],
85            num_rows: 0,
86            truncated_rows,
87            truncated_row_count: 0,
88        }
89    }
90
91    /// Decodes records from `input` returning the number of records and bytes read
92    ///
93    /// Note: this expects to be called with an empty `input` to signal EOF
94    pub fn decode(&mut self, input: &[u8], to_read: usize) -> Result<(usize, usize), ArrowError> {
95        if to_read == 0 {
96            return Ok((0, 0));
97        }
98
99        // Reserve sufficient capacity in offsets
100        self.offsets
101            .resize(self.offsets_len + to_read * self.num_columns, 0);
102
103        // The current offset into `input`
104        let mut input_offset = 0;
105
106        // The number of rows decoded in this pass
107        let mut read = 0;
108
109        loop {
110            // Reserve necessary space in output data based on best estimate
111            let remaining_rows = to_read - read;
112            let capacity = remaining_rows * self.num_columns * AVERAGE_FIELD_SIZE;
113            let estimated_data = capacity.max(MIN_CAPACITY);
114            self.data.resize(self.data_len + estimated_data, 0);
115
116            // Try to read a record
117            loop {
118                let (result, bytes_read, bytes_written, end_positions) =
119                    self.delimiter.read_record(
120                        &input[input_offset..],
121                        &mut self.data[self.data_len..],
122                        &mut self.offsets[self.offsets_len..],
123                    );
124
125                self.current_field += end_positions;
126                self.offsets_len += end_positions;
127                input_offset += bytes_read;
128                self.data_len += bytes_written;
129
130                match result {
131                    ReadRecordResult::End | ReadRecordResult::InputEmpty => {
132                        // Reached end of input
133                        return Ok((read, input_offset));
134                    }
135                    // Need to allocate more capacity
136                    ReadRecordResult::OutputFull => break,
137                    ReadRecordResult::OutputEndsFull => {
138                        return Err(ArrowError::CsvError(format!(
139                            "incorrect number of fields for line {}, expected {} got more than {}",
140                            self.line_number, self.num_columns, self.current_field
141                        )));
142                    }
143                    ReadRecordResult::Record => {
144                        if self.current_field != self.num_columns {
145                            if self.truncated_rows && self.current_field < self.num_columns {
146                                // If the number of fields is less than expected, pad with nulls
147                                let fill_count = self.num_columns - self.current_field;
148                                let fill_value = self.offsets[self.offsets_len - 1];
149                                self.offsets[self.offsets_len..self.offsets_len + fill_count]
150                                    .fill(fill_value);
151                                self.offsets_len += fill_count;
152                                self.truncated_row_count += 1;
153                            } else {
154                                return Err(ArrowError::CsvError(format!(
155                                    "incorrect number of fields for line {}, expected {} got {}",
156                                    self.line_number, self.num_columns, self.current_field
157                                )));
158                            }
159                        }
160                        read += 1;
161                        self.current_field = 0;
162                        self.line_number += 1;
163                        self.num_rows += 1;
164
165                        if read == to_read {
166                            // Read sufficient rows
167                            return Ok((read, input_offset));
168                        }
169
170                        if input.len() == input_offset {
171                            // Input exhausted, need to read more
172                            // Without this read_record will interpret the empty input
173                            // byte array as indicating the end of the file
174                            return Ok((read, input_offset));
175                        }
176                    }
177                }
178            }
179        }
180    }
181
182    /// Returns the current number of buffered records
183    pub fn len(&self) -> usize {
184        self.num_rows
185    }
186
187    /// Returns true if the decoder is empty
188    pub fn is_empty(&self) -> bool {
189        self.num_rows == 0
190    }
191
192    /// Returns the number of rows padded because they had fewer fields than expected
193    ///
194    /// Cumulative across [`Self::flush`] calls, reset by [`Self::clear`]
195    pub fn truncated_row_count(&self) -> usize {
196        self.truncated_row_count
197    }
198
199    /// Clears the current contents of the decoder
200    pub fn clear(&mut self) {
201        // This does not reset current_field to allow clearing part way through a record
202        self.offsets_len = 1;
203        self.data_len = 0;
204        self.num_rows = 0;
205        // The rows counted so far are being discarded along with the buffered data,
206        // so they must not be reported as padded
207        self.truncated_row_count = 0;
208    }
209
210    /// Flushes the current contents of the reader
211    pub fn flush(&mut self) -> Result<StringRecords<'_>, ArrowError> {
212        if self.current_field != 0 {
213            return Err(ArrowError::CsvError(
214                "Cannot flush part way through record".to_string(),
215            ));
216        }
217
218        // csv_core::Reader writes end offsets relative to the start of the row
219        // Therefore scan through and offset these based on the cumulative row offsets
220        let mut row_offset: usize = 0;
221        self.offsets[1..self.offsets_len]
222            .chunks_exact_mut(self.num_columns)
223            .try_for_each(|row| -> Result<(), ArrowError> {
224                let offset = row_offset;
225                row.iter_mut().try_for_each(|x| -> Result<(), ArrowError> {
226                    *x = x.checked_add(offset).ok_or_else(|| {
227                        ArrowError::CsvError(
228                            "CSV record offsets overflowed usize while flushing".to_string(),
229                        )
230                    })?;
231                    row_offset = *x;
232                    Ok(())
233                })
234            })?;
235
236        // Need to truncate data t1o the actual amount of data read
237        let data = std::str::from_utf8(&self.data[..self.data_len]).map_err(|e| {
238            let valid_up_to = e.valid_up_to();
239
240            // We can't use binary search because of empty fields
241            let idx = self.offsets[..self.offsets_len]
242                .iter()
243                .rposition(|x| *x <= valid_up_to)
244                .unwrap();
245
246            let field = idx % self.num_columns + 1;
247            let line_offset = self.line_number - self.num_rows;
248            let line = line_offset + idx / self.num_columns;
249
250            ArrowError::CsvError(format!(
251                "Encountered invalid UTF-8 data for line {line} and field {field}"
252            ))
253        })?;
254
255        let offsets = &self.offsets[..self.offsets_len];
256        let num_rows = self.num_rows;
257
258        // Reset state
259        // `truncated_row_count` is deliberately left alone so that it accumulates
260        // across the batches produced by a single decoder
261        self.offsets_len = 1;
262        self.data_len = 0;
263        self.num_rows = 0;
264
265        Ok(StringRecords {
266            num_rows,
267            num_columns: self.num_columns,
268            offsets,
269            data,
270        })
271    }
272}
273
274/// A collection of parsed, UTF-8 CSV records
275#[derive(Debug)]
276pub struct StringRecords<'a> {
277    num_columns: usize,
278    num_rows: usize,
279    offsets: &'a [usize],
280    data: &'a str,
281}
282
283impl<'a> StringRecords<'a> {
284    fn get(&self, index: usize) -> StringRecord<'a> {
285        let field_idx = index * self.num_columns;
286        StringRecord {
287            data: self.data,
288            offsets: &self.offsets[field_idx..field_idx + self.num_columns + 1],
289        }
290    }
291
292    pub fn len(&self) -> usize {
293        self.num_rows
294    }
295
296    pub fn iter(&self) -> impl Iterator<Item = StringRecord<'a>> + '_ {
297        (0..self.num_rows).map(|x| self.get(x))
298    }
299}
300
301/// A single parsed, UTF-8 CSV record
302#[derive(Debug, Clone, Copy)]
303pub struct StringRecord<'a> {
304    data: &'a str,
305    offsets: &'a [usize],
306}
307
308impl<'a> StringRecord<'a> {
309    pub fn get(&self, index: usize) -> &'a str {
310        let end = self.offsets[index + 1];
311        let start = self.offsets[index];
312
313        // SAFETY:
314        // Parsing produces offsets at valid byte boundaries
315        unsafe { self.data.get_unchecked(start..end) }
316    }
317}
318
319impl std::fmt::Display for StringRecord<'_> {
320    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
321        let num_fields = self.offsets.len() - 1;
322        write!(f, "[")?;
323        for i in 0..num_fields {
324            if i > 0 {
325                write!(f, ",")?;
326            }
327            write!(f, "{}", self.get(i))?;
328        }
329        write!(f, "]")?;
330        Ok(())
331    }
332}
333
334#[cfg(test)]
335mod tests {
336    use crate::reader::records::RecordDecoder;
337    use csv_core::Reader;
338    use std::io::{BufRead, BufReader, Cursor};
339
340    #[test]
341    fn test_basic() {
342        let csv = [
343            "foo,bar,baz",
344            "a,b,c",
345            "12,3,5",
346            "\"asda\"\"asas\",\"sdffsnsd\", as",
347        ]
348        .join("\n");
349
350        let mut expected = vec![
351            vec!["foo", "bar", "baz"],
352            vec!["a", "b", "c"],
353            vec!["12", "3", "5"],
354            vec!["asda\"asas", "sdffsnsd", " as"],
355        ]
356        .into_iter();
357
358        let mut reader = BufReader::with_capacity(3, Cursor::new(csv.as_bytes()));
359        let mut decoder = RecordDecoder::new(Reader::new(), 3, false);
360
361        loop {
362            let to_read = 3;
363            let mut read = 0;
364            loop {
365                let buf = reader.fill_buf().unwrap();
366                let (records, bytes) = decoder.decode(buf, to_read - read).unwrap();
367
368                reader.consume(bytes);
369                read += records;
370
371                if read == to_read || bytes == 0 {
372                    break;
373                }
374            }
375            if read == 0 {
376                break;
377            }
378
379            let b = decoder.flush().unwrap();
380            b.iter().zip(&mut expected).for_each(|(record, expected)| {
381                let actual = (0..3)
382                    .map(|field_idx| record.get(field_idx))
383                    .collect::<Vec<_>>();
384                assert_eq!(actual, expected)
385            });
386        }
387        assert!(expected.next().is_none());
388    }
389
390    #[test]
391    fn test_invalid_fields() {
392        let csv = "a,b\nb,c\na\n";
393        let mut decoder = RecordDecoder::new(Reader::new(), 2, false);
394        let err = decoder.decode(csv.as_bytes(), 4).unwrap_err().to_string();
395
396        let expected = "Csv error: incorrect number of fields for line 3, expected 2 got 1";
397
398        assert_eq!(err, expected);
399
400        // Test with initial skip
401        let mut decoder = RecordDecoder::new(Reader::new(), 2, false);
402        let (skipped, bytes) = decoder.decode(csv.as_bytes(), 1).unwrap();
403        assert_eq!(skipped, 1);
404        decoder.clear();
405
406        let remaining = &csv.as_bytes()[bytes..];
407        let err = decoder.decode(remaining, 3).unwrap_err().to_string();
408        assert_eq!(err, expected);
409    }
410
411    #[test]
412    fn test_skip_insufficient_rows() {
413        let csv = "a\nv\n";
414        let mut decoder = RecordDecoder::new(Reader::new(), 1, false);
415        let (read, bytes) = decoder.decode(csv.as_bytes(), 3).unwrap();
416        assert_eq!(read, 2);
417        assert_eq!(bytes, csv.len());
418    }
419
420    #[test]
421    fn test_truncated_rows() {
422        let csv = "a,b\nv\n,1\n,2\n,3\n";
423        let mut decoder = RecordDecoder::new(Reader::new(), 2, true);
424        let (read, bytes) = decoder.decode(csv.as_bytes(), 5).unwrap();
425        assert_eq!(read, 5);
426        assert_eq!(bytes, csv.len());
427        // Only "v" is short, the rows starting with a delimiter have both fields
428        assert_eq!(decoder.truncated_row_count(), 1);
429    }
430
431    #[test]
432    fn test_truncated_row_count_not_reset_by_flush() {
433        let csv = "a\nb,2\nc\n";
434        let mut decoder = RecordDecoder::new(Reader::new(), 2, true);
435
436        let (read, _) = decoder.decode(csv.as_bytes(), 2).unwrap();
437        assert_eq!(read, 2);
438        assert_eq!(decoder.truncated_row_count(), 1);
439        decoder.flush().unwrap();
440        assert_eq!(decoder.truncated_row_count(), 1);
441
442        let (read, _) = decoder.decode(&csv.as_bytes()[6..], 1).unwrap();
443        assert_eq!(read, 1);
444        assert_eq!(decoder.truncated_row_count(), 2);
445        decoder.flush().unwrap();
446        assert_eq!(decoder.truncated_row_count(), 2);
447    }
448
449    #[test]
450    fn test_truncated_row_count_reset_by_clear() {
451        let csv = "a\nb,2\n";
452        let mut decoder = RecordDecoder::new(Reader::new(), 2, true);
453
454        let (read, _) = decoder.decode(csv.as_bytes(), 2).unwrap();
455        assert_eq!(read, 2);
456        assert_eq!(decoder.truncated_row_count(), 1);
457
458        // The rows are discarded, so the padding done to them is discarded too
459        decoder.clear();
460        assert_eq!(decoder.truncated_row_count(), 0);
461    }
462
463    /// Regression test for an overflow path found by the `arrow-csv`
464    /// cargo-fuzz harness being prototyped for #5332. Stages the
465    /// `RecordDecoder` state directly so that rebasing the second row's
466    /// end offset overflows `usize`. With the previous `*x += offset` this
467    /// panicked with `attempt to add with overflow`; the patched code
468    /// surfaces the condition as `ArrowError::CsvError`.
469    #[test]
470    fn test_flush_offset_overflow_returns_csv_error() {
471        let mut decoder = RecordDecoder::new(Reader::new(), 1, false);
472        decoder.offsets = vec![0, usize::MAX, 1];
473        decoder.offsets_len = 3;
474        decoder.num_rows = 2;
475        let err = decoder.flush().unwrap_err();
476        assert_eq!(
477            err.to_string(),
478            "Csv error: CSV record offsets overflowed usize while flushing"
479        );
480    }
481}