diff --git a/arrow-csv/src/lib.rs b/arrow-csv/src/lib.rs index ea70f205f7e5..12776f86eaa7 100644 --- a/arrow-csv/src/lib.rs +++ b/arrow-csv/src/lib.rs @@ -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; diff --git a/arrow-csv/src/reader/mod.rs b/arrow-csv/src/reader/mod.rs index 56f916512936..f5af8371145a 100644 --- a/arrow-csv/src/reader/mod.rs +++ b/arrow-csv/src/reader/mod.rs @@ -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`] diff --git a/arrow-csv/src/reader/records.rs b/arrow-csv/src/reader/records.rs index f93e365ca6f9..7db67ee905e9 100644 --- a/arrow-csv/src/reader/records.rs +++ b/arrow-csv/src/reader/records.rs @@ -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; @@ -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]; @@ -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}; @@ -507,12 +513,22 @@ mod tests { record: Vec, } + #[derive(Debug, Clone, Eq, PartialEq)] + struct OwnedRecord { + line_number: usize, + byte_offset: usize, + record: Vec, + } + #[derive(Debug, Default)] - struct CollectRecordErrors(Mutex>); + struct CollectRecords { + errors: Mutex>, + records: Mutex>, + } - 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, @@ -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] @@ -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)); @@ -624,7 +649,7 @@ mod tests { ] ); - let errors = handler.0.lock().unwrap(); + let errors = handler.errors.lock().unwrap(); assert_eq!( *errors, [ @@ -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] @@ -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;