diff --git a/datafusion/datasource-csv/src/file_format.rs b/datafusion/datasource-csv/src/file_format.rs index 881543266bd2..1ad1a71561e4 100644 --- a/datafusion/datasource-csv/src/file_format.rs +++ b/datafusion/datasource-csv/src/file_format.rs @@ -21,7 +21,7 @@ use std::collections::{HashMap, HashSet}; use std::fmt::{self, Debug}; use std::sync::Arc; -use crate::source::CsvSource; +use crate::source::{CsvRecordErrorHandlerFactory, CsvSource}; use arrow::array::RecordBatch; use arrow::csv::reader::SnowflakeCsvTypeState; @@ -152,6 +152,7 @@ impl GetExt for CsvFormatFactory { pub struct CsvFormat { options: CsvOptions, record_error_handler: Option>, + record_error_handler_factory: Option>, numeric_boolean_values: bool, } @@ -235,6 +236,18 @@ impl CsvFormat { handler: Arc, ) -> Self { self.record_error_handler = Some(handler); + self.record_error_handler_factory = None; + self + } + + /// Use a file-specific handler for malformed and decoded CSV records. + /// The factory is called when each source file is opened. + pub fn with_record_error_handler_factory( + mut self, + factory: Arc, + ) -> Self { + self.record_error_handler = None; + self.record_error_handler_factory = Some(factory); self } @@ -563,6 +576,8 @@ impl FileFormat for CsvFormat { .with_numeric_boolean_values(self.numeric_boolean_values); if let Some(handler) = &self.record_error_handler { source = source.with_record_error_handler(Arc::clone(handler)); + } else if let Some(factory) = &self.record_error_handler_factory { + source = source.with_record_error_handler_factory(Arc::clone(factory)); } Arc::new(source) } diff --git a/datafusion/datasource-csv/src/source.rs b/datafusion/datasource-csv/src/source.rs index eb3411a83400..c1dc13f3ebb4 100644 --- a/datafusion/datasource-csv/src/source.rs +++ b/datafusion/datasource-csv/src/source.rs @@ -47,10 +47,18 @@ use datafusion_physical_plan::{ use crate::file_format::CsvDecoder; use futures::{StreamExt, TryStreamExt}; use object_store::buffered::BufWriter; +use object_store::path::Path; use object_store::{GetOptions, GetResultPayload, ObjectStore}; use regex::Regex; use tokio::io::AsyncWriteExt; +/// Creates a CSV record handler scoped to the file being opened. +pub trait CsvRecordErrorHandlerFactory: fmt::Debug + Send + Sync { + /// Return the handler for `location`. The location is the object-store path + /// of one file, before its CSV reader is created. + fn for_file(&self, location: &Path) -> Arc; +} + /// A Config for [`CsvOpener`] /// /// # Example: create a `DataSourceExec` for CSV @@ -89,6 +97,7 @@ pub struct CsvSource { options: CsvOptions, numeric_boolean_values: bool, record_error_handler: Option>, + record_error_handler_factory: Option>, batch_size: Option, table_schema: TableSchema, projection: SplitProjection, @@ -103,6 +112,7 @@ impl CsvSource { options: CsvOptions::default(), numeric_boolean_values: false, record_error_handler: None, + record_error_handler_factory: None, projection: SplitProjection::unprojected(&table_schema), table_schema, batch_size: None, @@ -134,6 +144,18 @@ impl CsvSource { handler: Arc, ) -> Self { self.record_error_handler = Some(handler); + self.record_error_handler_factory = None; + self + } + + /// Select a record handler for each opened file. This disables byte-range + /// repartitioning, as record offsets are file-relative. + pub fn with_record_error_handler_factory( + mut self, + factory: Arc, + ) -> Self { + self.record_error_handler = None; + self.record_error_handler_factory = Some(factory); self } @@ -353,6 +375,7 @@ impl FileSource for CsvSource { // Cannot repartition if values may contain newlines, as record // boundaries cannot be determined by byte offset alone self.record_error_handler.is_none() + && self.record_error_handler_factory.is_none() && !self.options.newlines_in_values.unwrap_or(false) } @@ -384,7 +407,9 @@ impl FileSource for CsvSource { use datafusion_proto_models::protobuf; use protobuf::physical_plan_node::PhysicalPlanType; - if self.record_error_handler.is_some() { + if self.record_error_handler.is_some() + || self.record_error_handler_factory.is_some() + { return datafusion_common::not_impl_err!( "CSV record error handlers cannot be serialized" ); @@ -463,6 +488,10 @@ impl FileOpener for CsvOpener { } let mut config = (*self.config).clone(); + if let Some(factory) = &config.record_error_handler_factory { + config.record_error_handler = + Some(factory.for_file(&partitioned_file.object_meta.location)); + } config.options.has_header = Some(csv_has_header); config.options.truncated_rows = Some(config.truncate_rows()); @@ -579,7 +608,7 @@ pub async fn plan_to_csv( let storeref = Arc::clone(&store); let plan: Arc = Arc::clone(&plan); let filename = format!("{}/part-{i}.csv", parsed.prefix()); - let file = object_store::path::Path::parse(filename)?; + let file = Path::parse(filename)?; let mut stream = plan.execute(i, Arc::clone(&task_ctx))?; join_set.spawn(async move { @@ -723,6 +752,9 @@ mod tests { use arrow::csv::{CsvRecordError, CsvRecordErrorHandler}; use arrow::datatypes::{DataType, Field, Schema}; use arrow::error::ArrowError; + use bytes::Bytes; + use object_store::ObjectStoreExt; + use object_store::memory::InMemory; use std::io::Cursor; use std::sync::Mutex; @@ -848,6 +880,97 @@ mod tests { } } + #[derive(Debug)] + struct TaggedRecords { + path: String, + seen: Arc>>, + } + + impl CsvRecordErrorHandler for TaggedRecords { + fn handle(&self, error: &CsvRecordError<'_>) -> Result<(), ArrowError> { + self.seen.lock().unwrap().push(( + self.path.clone(), + format!("error:{}", String::from_utf8_lossy(error.record)), + )); + Ok(()) + } + + fn handle_record(&self, record: &csv::CsvRecord<'_>) -> Result<(), ArrowError> { + self.seen.lock().unwrap().push(( + self.path.clone(), + String::from_utf8_lossy(record.record).into_owned(), + )); + Ok(()) + } + } + + #[derive(Debug)] + struct TaggedRecordsFactory { + seen: Arc>>, + } + + impl CsvRecordErrorHandlerFactory for TaggedRecordsFactory { + fn for_file(&self, location: &Path) -> Arc { + Arc::new(TaggedRecords { + path: location.to_string(), + seen: Arc::clone(&self.seen), + }) + } + } + + #[tokio::test] + async fn record_error_handler_factory_is_scoped_to_opened_file() + -> Result<(), Box> { + let store = Arc::new(InMemory::new()); + let seen = Arc::new(Mutex::new(Vec::new())); + let mut source = CsvSource::new(Arc::new(Schema::new(vec![Field::new( + "value", + DataType::Utf8, + false, + )]))) + .with_csv_options(CsvOptions { + has_header: Some(false), + ..Default::default() + }) + .with_record_error_handler_factory(Arc::new(TaggedRecordsFactory { + seen: Arc::clone(&seen), + })); + source.batch_size = Some(1024); + assert!(!source.supports_repartitioning()); + let opener = CsvOpener::new( + Arc::new(source), + FileCompressionType::UNCOMPRESSED, + store.clone(), + ); + for (name, content) in [ + ("first.csv", "first\nbad,extra\n"), + ("second.csv", "second\nwrong,extra\n"), + ] { + let path = Path::from(name); + store.put(&path, Bytes::from(content).into()).await?; + let meta = store.head(&path).await?; + let batches = opener + .open(PartitionedFile::new_from_meta(meta))? + .await? + .try_collect::>() + .await?; + assert_eq!( + batches.iter().map(|batch| batch.num_rows()).sum::(), + 1 + ); + } + assert_eq!( + *seen.lock().unwrap(), + [ + ("first.csv".to_owned(), "first\n".to_owned()), + ("first.csv".to_owned(), "error:bad,extra\n".to_owned()), + ("second.csv".to_owned(), "second\n".to_owned()), + ("second.csv".to_owned(), "error:wrong,extra\n".to_owned()) + ] + ); + Ok(()) + } + #[test] fn record_error_handler_skips_malformed_rows_and_disables_repartitioning() { let schema = Arc::new(Schema::new(vec![