diff --git a/datafusion/datasource-csv/Cargo.toml b/datafusion/datasource-csv/Cargo.toml index 42200d67c57a..43bb666b5b4a 100644 --- a/datafusion/datasource-csv/Cargo.toml +++ b/datafusion/datasource-csv/Cargo.toml @@ -57,6 +57,9 @@ regex = { workspace = true } regex-syntax = "0.8" tokio = { workspace = true } +[dev-dependencies] +datafusion-datasource = { workspace = true, features = ["compression"] } + # Note: add additional linter rules in lib.rs. # Rust does not support workspace + new linter rules in subcrates yet # https://github.com/rust-lang/cargo/issues/13157 diff --git a/datafusion/datasource-csv/src/file_format.rs b/datafusion/datasource-csv/src/file_format.rs index 8943897dc8ec..881543266bd2 100644 --- a/datafusion/datasource-csv/src/file_format.rs +++ b/datafusion/datasource-csv/src/file_format.rs @@ -65,6 +65,22 @@ use object_store::{ }; use regex::Regex; +async fn is_uncompressed_auto_file(store: &dyn ObjectStore, file: &ObjectMeta) -> bool { + if file.size == 0 { + return true; + } + let len = file + .size + .min(u64::try_from(FileCompressionType::AUTO_PROBE_BYTES).unwrap_or(u64::MAX)); + store + .get_range(&file.location, 0..len) + .await + .is_ok_and(|header| { + FileCompressionType::detect_from_header(&header) + == FileCompressionType::UNCOMPRESSED + }) +} + #[derive(Default)] /// Factory used to create [`CsvFormat`] pub struct CsvFormatFactory { @@ -482,17 +498,7 @@ impl FileFormat for CsvFormat { let plain = futures::future::join_all(files.into_iter().map(|file| { let store = Arc::clone(&store); async move { - if file.object_meta.size == 0 { - return true; - } - let len = file.object_meta.size.min(6); - store - .get_range(&file.object_meta.location, 0..len) - .await - .is_ok_and(|header| { - FileCompressionType::detect_from_header(&header) - == FileCompressionType::UNCOMPRESSED - }) + is_uncompressed_auto_file(store.as_ref(), &file.object_meta).await } })) .await; @@ -1273,11 +1279,12 @@ impl From<&CsvFormatFactory> for datafusion_proto_models::protobuf::CsvOptions { #[cfg(test)] mod tests { - use super::{CsvFormat, build_schema_helper}; + use super::{CsvFormat, build_schema_helper, is_uncompressed_auto_file}; use arrow::datatypes::DataType; use bytes::Bytes; use datafusion_common::DFSchema; use datafusion_common::config::TableOptions; + use datafusion_datasource::file_compression_type::FileCompressionType; use datafusion_execution::TaskContext; use datafusion_execution::config::SessionConfig; use datafusion_execution::runtime_env::RuntimeEnv; @@ -1289,11 +1296,49 @@ mod tests { use datafusion_physical_expr_common::physical_expr::PhysicalExpr; use datafusion_physical_plan::ExecutionPlan; use datafusion_session::{CatalogProviderList, EmptyCatalogProviderList, Session}; - use futures::stream; + use futures::{StreamExt, TryStreamExt, stream}; + use object_store::ObjectStoreExt; + use object_store::memory::InMemory; + use object_store::path::Path; use std::any::Any; use std::collections::{HashMap, HashSet}; use std::sync::Arc; + #[tokio::test] + async fn auto_repartition_probe_checks_past_printable_zlib_header() + -> Result<(), Box> { + assert!(FileCompressionType::compression_enabled()); + let store = InMemory::new(); + // zlib.compress(b"1000,payload\n", level=2): its first six bytes are printable. + let zlib: &[u8] = b"x^3400\xd0)H\xac\xcc\xc9OL\xe1\x02\x00\x19\x10\x03\xe2"; + let decoded = FileCompressionType::DEFLATE + .convert_stream(stream::once(async { Ok(Bytes::from_static(zlib)) }).boxed())? + .try_collect::>() + .await? + .concat(); + assert_eq!(decoded, b"1000,payload\n"); + assert_eq!( + FileCompressionType::detect_from_header(&zlib[..6]), + FileCompressionType::UNCOMPRESSED + ); + for (name, content, expected_plain) in [ + ("compressed.csv", zlib, false), + ("plain.csv", b"80,payload-0\n".as_slice(), true), + ] { + let path = Path::from(name); + store + .put(&path, Bytes::copy_from_slice(content).into()) + .await?; + let metadata = store.head(&path).await?; + assert_eq!( + is_uncompressed_auto_file(&store, &metadata).await, + expected_plain, + "{name}" + ); + } + Ok(()) + } + struct InferenceSession { config: SessionConfig, runtime_env: Arc, diff --git a/datafusion/datasource/src/file_compression_type.rs b/datafusion/datasource/src/file_compression_type.rs index 8ae2d6a4c6bc..62aa7da7ae13 100644 --- a/datafusion/datasource/src/file_compression_type.rs +++ b/datafusion/datasource/src/file_compression_type.rs @@ -48,6 +48,8 @@ use bytes::Bytes; use bzip2::read::MultiBzDecoder; #[cfg(feature = "compression")] use flate2::read::{DeflateDecoder, MultiGzDecoder, ZlibDecoder}; +#[cfg(feature = "compression")] +use flate2::{Decompress, FlushDecompress}; use futures::StreamExt; #[cfg(feature = "compression")] use futures::TryStreamExt; @@ -108,6 +110,9 @@ impl FromStr for FileCompressionType { /// `FileCompressionType` implementation impl FileCompressionType { + /// Header bytes inspected for ambiguous zlib AUTO detection. + pub const AUTO_PROBE_BYTES: usize = 64; + /// Whether this build can decode compressed input or use AUTO detection. pub const fn compression_enabled() -> bool { cfg!(feature = "compression") @@ -346,10 +351,10 @@ impl FileCompressionType { #[cfg(feature = "compression")] AUTO => futures::stream::once(async move { let mut source = s; - let mut header = [0_u8; 6]; + let mut header = [0_u8; Self::AUTO_PROBE_BYTES]; let mut header_len = 0; let mut initial_chunks = Vec::new(); - while header_len < header.len() { + while header_len < AUTO_HEADER_BYTES { match source.next().await { Some(Ok(bytes)) if bytes.is_empty() => continue, Some(Ok(bytes)) => { @@ -363,6 +368,22 @@ impl FileCompressionType { None => break, } } + if Self::is_zlib_header(&header[..header_len]) { + while header_len < header.len() { + match source.next().await { + Some(Ok(bytes)) if bytes.is_empty() => continue, + Some(Ok(bytes)) => { + let copied = (header.len() - header_len).min(bytes.len()); + header[header_len..header_len + copied] + .copy_from_slice(&bytes[..copied]); + header_len += copied; + initial_chunks.push(Ok(bytes)); + } + Some(Err(error)) => return Err(error), + None => break, + } + } + } let replay = futures::stream::iter(initial_chunks).chain(source).boxed(); Self::detect_from_header(&header[..header_len]).convert_stream(replay) }) @@ -410,15 +431,24 @@ impl FileCompressionType { #[cfg(feature = "compression")] AUTO => { let mut source = r; - let mut header = [0_u8; 6]; + let mut header = [0_u8; Self::AUTO_PROBE_BYTES]; let mut header_len = 0; - while header_len < header.len() { - let read = source.read(&mut header[header_len..])?; + while header_len < AUTO_HEADER_BYTES { + let read = source.read(&mut header[header_len..AUTO_HEADER_BYTES])?; if read == 0 { break; } header_len += read; } + if Self::is_zlib_header(&header[..header_len]) { + while header_len < header.len() { + let read = source.read(&mut header[header_len..])?; + if read == 0 { + break; + } + header_len += read; + } + } let replay: Box = Box::new( std::io::Cursor::new(header[..header_len].to_vec()).chain(source), ); @@ -434,9 +464,14 @@ impl FileCompressionType { } } +#[cfg(feature = "compression")] +const AUTO_HEADER_BYTES: usize = 6; impl FileCompressionType { - /// Identify codecs supported by AUTO from up to the first six file bytes. + /// Identify codecs supported by AUTO from a bounded file prefix. /// Headerless codecs such as Brotli and raw Deflate require explicit selection. + /// Zlib has no unique magic bytes: printable text takes precedence when its + /// header happens to satisfy the RFC 1950 checksum. Explicit DEFLATE avoids + /// that ambiguity when the compressed prefix itself resembles plain text. pub fn detect_from_header(header: &[u8]) -> Self { if header.starts_with(&[0x1f, 0x8b]) { Self::GZIP @@ -446,16 +481,48 @@ impl FileCompressionType { Self::ZSTD } else if header.starts_with(&[0xfd, b'7', b'z', b'X', b'Z', 0x00]) { Self::XZ - } else if header.len() >= 2 - && header[0] & 0x0f == 8 - && header[0] >> 4 <= 7 - && u16::from_be_bytes([header[0], header[1]]).is_multiple_of(31) + } else if Self::is_zlib_header(header) + && !Self::looks_like_plain_text(header) + && Self::zlib_prefix_decodes(header) { Self::DEFLATE } else { Self::UNCOMPRESSED } } + + fn is_zlib_header(header: &[u8]) -> bool { + header.len() >= 2 + && header[0] & 0x0f == 8 + && header[0] >> 4 <= 7 + && u16::from_be_bytes([header[0], header[1]]).is_multiple_of(31) + } + + #[cfg(feature = "compression")] + fn looks_like_plain_text(header: &[u8]) -> bool { + header.iter().all(|byte| { + byte.is_ascii_graphic() || matches!(*byte, b'\n' | b'\r' | b'\t' | b' ') + }) + } + + #[cfg(not(feature = "compression"))] + fn looks_like_plain_text(_header: &[u8]) -> bool { + false + } + + #[cfg(feature = "compression")] + fn zlib_prefix_decodes(header: &[u8]) -> bool { + let mut decoder = Decompress::new(true); + let mut output = [0_u8; 4096]; + decoder + .decompress(header, &mut output, FlushDecompress::None) + .is_ok() + } + + #[cfg(not(feature = "compression"))] + fn zlib_prefix_decodes(_header: &[u8]) -> bool { + true + } } /// Trait for extending the functionality of the `FileType` enum. @@ -559,6 +626,120 @@ mod tests { Ok(()) } + #[cfg(feature = "compression")] + #[tokio::test] + async fn auto_keeps_zlib_like_csv_headers_uncompressed() -> Result<(), DataFusionError> + { + use futures::TryStreamExt; + use std::io::Read; + + for plain in [ + b"80,payload-0\n".as_slice(), + b"x^,payload-0\n", + b"x^txex6Qs2qs_CZPGEGnB15C..mSVKfsXkgaEDaYi4i3MiXH,gWc0ieiSkcBOqWR", + ] { + assert_eq!( + FileCompressionType::detect_from_header(plain), + FileCompressionType::UNCOMPRESSED + ); + let chunks = plain + .chunks(1) + .map(|chunk| Ok(Bytes::copy_from_slice(chunk))) + .collect::>>(); + let decoded = FileCompressionType::AUTO + .convert_stream(futures::stream::iter(chunks).boxed())? + .try_collect::>() + .await? + .concat(); + assert_eq!(decoded, plain); + + let mut reader = + FileCompressionType::AUTO.convert_read(std::io::Cursor::new(plain))?; + let mut decoded = Vec::new(); + reader.read_to_end(&mut decoded)?; + assert_eq!(decoded, plain); + } + Ok(()) + } + + #[cfg(feature = "compression")] + #[test] + fn auto_read_only_probes_six_bytes_for_non_zlib_input() -> Result<(), DataFusionError> + { + use std::io::{Cursor, Read}; + use std::sync::Arc; + use std::sync::atomic::{AtomicUsize, Ordering}; + + struct CountedReader { + source: Cursor>, + bytes_read: Arc, + } + + impl Read for CountedReader { + fn read(&mut self, buf: &mut [u8]) -> std::io::Result { + let count = self.source.read(buf)?; + self.bytes_read.fetch_add(count, Ordering::Relaxed); + Ok(count) + } + } + + let plain = b"ordinary,uncompressed,csv\n"; + let bytes_read = Arc::new(AtomicUsize::new(0)); + let source = CountedReader { + source: Cursor::new(plain.to_vec()), + bytes_read: Arc::clone(&bytes_read), + }; + let mut decoded = FileCompressionType::AUTO.convert_read(source)?; + assert_eq!(bytes_read.load(Ordering::Relaxed), 6); + + let mut output = Vec::new(); + decoded.read_to_end(&mut output)?; + assert_eq!(output, plain); + Ok(()) + } + + #[cfg(feature = "compression")] + #[tokio::test] + async fn auto_preserves_truncated_binary_zlib_detection() + -> Result<(), DataFusionError> { + use futures::TryStreamExt; + use std::io::Read; + + let truncated = [0x78, 0x9c, 0x4b, 0x4c, 0x4a]; + assert_eq!( + FileCompressionType::detect_from_header(&truncated), + FileCompressionType::DEFLATE + ); + let chunks = truncated + .iter() + .map(|byte| Ok(Bytes::copy_from_slice(&[*byte]))) + .collect::>>(); + let stream_result = FileCompressionType::AUTO + .convert_stream(futures::stream::iter(chunks).boxed())? + .try_collect::>() + .await; + assert!(match stream_result { + Err(_) => true, + Ok(chunks) => chunks.concat() != truncated, + }); + + let mut reader = + FileCompressionType::AUTO.convert_read(std::io::Cursor::new(truncated))?; + let mut output = Vec::new(); + let read_result = reader.read_to_end(&mut output); + assert!(read_result.is_err() || output != truncated); + Ok(()) + } + + #[cfg(not(feature = "compression"))] + #[test] + fn auto_header_detection_is_unchanged_without_compression() { + assert_eq!( + FileCompressionType::detect_from_header(b"80,payload-0\n"), + FileCompressionType::DEFLATE + ); + } + #[cfg(feature = "compression")] #[tokio::test] async fn explicit_brotli_and_deflate_roundtrip() -> Result<(), DataFusionError> {