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
17 changes: 16 additions & 1 deletion datafusion/datasource-csv/src/file_format.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -152,6 +152,7 @@ impl GetExt for CsvFormatFactory {
pub struct CsvFormat {
options: CsvOptions,
record_error_handler: Option<Arc<dyn CsvRecordErrorHandler>>,
record_error_handler_factory: Option<Arc<dyn CsvRecordErrorHandlerFactory>>,
numeric_boolean_values: bool,
}

Expand Down Expand Up @@ -235,6 +236,18 @@ impl CsvFormat {
handler: Arc<dyn CsvRecordErrorHandler>,
) -> 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<dyn CsvRecordErrorHandlerFactory>,
) -> Self {
self.record_error_handler = None;
self.record_error_handler_factory = Some(factory);
self
}

Expand Down Expand Up @@ -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)
}
Expand Down
127 changes: 125 additions & 2 deletions datafusion/datasource-csv/src/source.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<dyn csv::CsvRecordErrorHandler>;
}

/// A Config for [`CsvOpener`]
///
/// # Example: create a `DataSourceExec` for CSV
Expand Down Expand Up @@ -89,6 +97,7 @@ pub struct CsvSource {
options: CsvOptions,
numeric_boolean_values: bool,
record_error_handler: Option<Arc<dyn csv::CsvRecordErrorHandler>>,
record_error_handler_factory: Option<Arc<dyn CsvRecordErrorHandlerFactory>>,
batch_size: Option<usize>,
table_schema: TableSchema,
projection: SplitProjection,
Expand All @@ -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,
Expand Down Expand Up @@ -134,6 +144,18 @@ impl CsvSource {
handler: Arc<dyn csv::CsvRecordErrorHandler>,
) -> 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<dyn CsvRecordErrorHandlerFactory>,
) -> Self {
self.record_error_handler = None;
self.record_error_handler_factory = Some(factory);
self
}

Expand Down Expand Up @@ -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)
}

Expand Down Expand Up @@ -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"
);
Expand Down Expand Up @@ -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());

Expand Down Expand Up @@ -579,7 +608,7 @@ pub async fn plan_to_csv(
let storeref = Arc::clone(&store);
let plan: Arc<dyn ExecutionPlan> = 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 {
Expand Down Expand Up @@ -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;

Expand Down Expand Up @@ -848,6 +880,97 @@ mod tests {
}
}

#[derive(Debug)]
struct TaggedRecords {
path: String,
seen: Arc<Mutex<Vec<(String, String)>>>,
}

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<Mutex<Vec<(String, String)>>>,
}

impl CsvRecordErrorHandlerFactory for TaggedRecordsFactory {
fn for_file(&self, location: &Path) -> Arc<dyn CsvRecordErrorHandler> {
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<dyn std::error::Error>> {
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::<Vec<_>>()
.await?;
assert_eq!(
batches.iter().map(|batch| batch.num_rows()).sum::<usize>(),
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![
Expand Down
Loading