Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion arrow-csv/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,8 @@ pub mod reader;
pub mod writer;

pub use self::reader::{
CsvRecordError, CsvRecordErrorHandler, Reader, ReaderBuilder, infer_schema_from_files,
CsvRecord, CsvRecordError, CsvRecordErrorHandler, Reader, ReaderBuilder,
infer_schema_from_files,
};
pub use self::writer::QuoteStyle;
pub use self::writer::Writer;
Expand Down
19 changes: 19 additions & 0 deletions arrow-csv/src/reader/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -197,10 +197,29 @@ pub struct CsvRecordError<'a> {
pub record: &'a [u8],
}

/// Source metadata for a successfully decoded CSV record.
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
pub struct CsvRecord<'a> {
/// One-based record number, including any header row.
pub line_number: usize,
/// Zero-based byte offset of the start of the record in the input stream.
pub byte_offset: usize,
/// Original record bytes, including its record terminator when present.
pub record: &'a [u8],
}

/// Receives malformed CSV records that should be skipped instead of aborting the scan.
pub trait CsvRecordErrorHandler: Debug + Send + Sync {
/// Handle one malformed record. Returning an error aborts the scan.
fn handle(&self, error: &CsvRecordError<'_>) -> Result<(), ArrowError>;

/// Observe one successfully decoded record.
///
/// This callback is invoked only when a record error handler is configured. The default
/// implementation preserves the existing malformed-record-only behavior.
fn handle_record(&self, _record: &CsvRecord<'_>) -> Result<(), ArrowError> {
Ok(())
}
}

/// Order should match [`InferredDataType`]
Expand Down
59 changes: 51 additions & 8 deletions arrow-csv/src/reader/records.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ use arrow_schema::ArrowError;
use csv_core::{ReadRecordResult, Reader};
use std::sync::Arc;

use super::{CsvRecordError, CsvRecordErrorHandler};
use super::{CsvRecord, CsvRecordError, CsvRecordErrorHandler};

/// The estimated length of a field in bytes
const AVERAGE_FIELD_SIZE: usize = 8;
Expand Down Expand Up @@ -310,6 +310,12 @@ impl RecordDecoder {
continue;
}

handler.handle_record(&CsvRecord {
line_number: self.line_number,
byte_offset: self.record_byte_offset,
record: &self.record_bytes,
})?;

if self.current_field < self.num_columns {
let fill_count = self.num_columns - self.current_field;
let fill_value = self.offsets[self.offsets_len - 1];
Expand Down Expand Up @@ -490,7 +496,7 @@ impl std::fmt::Display for StringRecord<'_> {

#[cfg(test)]
mod tests {
use crate::reader::{CsvRecordError, CsvRecordErrorHandler};
use crate::reader::{CsvRecord, CsvRecordError, CsvRecordErrorHandler};
use arrow_schema::ArrowError;
use csv_core::Reader;
use std::io::{BufRead, BufReader, Cursor};
Expand All @@ -507,12 +513,22 @@ mod tests {
record: Vec<u8>,
}

#[derive(Debug, Clone, Eq, PartialEq)]
struct OwnedRecord {
line_number: usize,
byte_offset: usize,
record: Vec<u8>,
}

#[derive(Debug, Default)]
struct CollectRecordErrors(Mutex<Vec<OwnedRecordError>>);
struct CollectRecords {
errors: Mutex<Vec<OwnedRecordError>>,
records: Mutex<Vec<OwnedRecord>>,
}

impl CsvRecordErrorHandler for CollectRecordErrors {
impl CsvRecordErrorHandler for CollectRecords {
fn handle(&self, error: &CsvRecordError<'_>) -> Result<(), ArrowError> {
self.0.lock().unwrap().push(OwnedRecordError {
self.errors.lock().unwrap().push(OwnedRecordError {
line_number: error.line_number,
byte_offset: error.byte_offset,
expected_fields: error.expected_fields,
Expand All @@ -521,6 +537,15 @@ mod tests {
});
Ok(())
}

fn handle_record(&self, record: &CsvRecord<'_>) -> Result<(), ArrowError> {
self.records.lock().unwrap().push(OwnedRecord {
line_number: record.line_number,
byte_offset: record.byte_offset,
record: record.record.to_vec(),
});
Ok(())
}
}

#[test]
Expand Down Expand Up @@ -597,7 +622,7 @@ mod tests {
#[test]
fn test_invalid_fields_handler_skips_records_across_input_chunks() {
let csv = b"1,ok\n2,extra,value\n3\n4,after\n";
let handler = Arc::new(CollectRecordErrors::default());
let handler = Arc::new(CollectRecords::default());
let mut decoder = RecordDecoder::new(Reader::new(), 2, false)
.with_record_error_handler(Some(handler.clone()));
let mut reader = BufReader::with_capacity(3, Cursor::new(csv));
Expand All @@ -624,7 +649,7 @@ mod tests {
]
);

let errors = handler.0.lock().unwrap();
let errors = handler.errors.lock().unwrap();
assert_eq!(
*errors,
[
Expand All @@ -644,6 +669,24 @@ mod tests {
},
]
);
drop(errors);

let records = handler.records.lock().unwrap();
assert_eq!(
*records,
[
OwnedRecord {
line_number: 1,
byte_offset: 0,
record: b"1,ok\n".to_vec(),
},
OwnedRecord {
line_number: 4,
byte_offset: 21,
record: b"4,after\n".to_vec(),
},
]
);
}

#[test]
Expand All @@ -654,7 +697,7 @@ mod tests {
}
csv.push(b'\n');

let handler = Arc::new(CollectRecordErrors::default());
let handler = Arc::new(CollectRecords::default());
let mut decoder =
RecordDecoder::new(Reader::new(), 2, false).with_record_error_handler(Some(handler));
let mut input_offset = 0_usize;
Expand Down
Loading