1use arrow_schema::ArrowError;
19use csv_core::{ReadRecordResult, Reader};
20
21const AVERAGE_FIELD_SIZE: usize = 8;
23
24const MIN_CAPACITY: usize = 1024;
26
27#[derive(Debug)]
29pub struct RecordDecoder {
30 delimiter: Reader,
31
32 num_columns: usize,
34
35 line_number: usize,
37
38 offsets: Vec<usize>,
40
41 offsets_len: usize,
45
46 current_field: usize,
48
49 num_rows: usize,
51
52 data: Vec<u8>,
54
55 data_len: usize,
59
60 truncated_rows: bool,
65
66 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, 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 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 self.offsets
101 .resize(self.offsets_len + to_read * self.num_columns, 0);
102
103 let mut input_offset = 0;
105
106 let mut read = 0;
108
109 loop {
110 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 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 return Ok((read, input_offset));
134 }
135 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 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 return Ok((read, input_offset));
168 }
169
170 if input.len() == input_offset {
171 return Ok((read, input_offset));
175 }
176 }
177 }
178 }
179 }
180 }
181
182 pub fn len(&self) -> usize {
184 self.num_rows
185 }
186
187 pub fn is_empty(&self) -> bool {
189 self.num_rows == 0
190 }
191
192 pub fn truncated_row_count(&self) -> usize {
196 self.truncated_row_count
197 }
198
199 pub fn clear(&mut self) {
201 self.offsets_len = 1;
203 self.data_len = 0;
204 self.num_rows = 0;
205 self.truncated_row_count = 0;
208 }
209
210 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 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 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 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 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#[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#[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 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 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 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 decoder.clear();
460 assert_eq!(decoder.truncated_row_count(), 0);
461 }
462
463 #[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}