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: 3 additions & 0 deletions datafusion/datasource-csv/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
71 changes: 58 additions & 13 deletions datafusion/datasource-csv/src/file_format.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -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;
Expand All @@ -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<dyn std::error::Error>> {
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::<Vec<_>>()
.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<RuntimeEnv>,
Expand Down
201 changes: 191 additions & 10 deletions datafusion/datasource/src/file_compression_type.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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)) => {
Expand All @@ -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)
})
Expand Down Expand Up @@ -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<dyn Read + Send> = Box::new(
std::io::Cursor::new(header[..header_len].to_vec()).chain(source),
);
Expand All @@ -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
Expand All @@ -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.
Expand Down Expand Up @@ -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::<Vec<Result<Bytes, DataFusionError>>>();
let decoded = FileCompressionType::AUTO
.convert_stream(futures::stream::iter(chunks).boxed())?
.try_collect::<Vec<_>>()
.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<Vec<u8>>,
bytes_read: Arc<AtomicUsize>,
}

impl Read for CountedReader {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
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::<Vec<Result<Bytes, DataFusionError>>>();
let stream_result = FileCompressionType::AUTO
.convert_stream(futures::stream::iter(chunks).boxed())?
.try_collect::<Vec<_>>()
.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> {
Expand Down
Loading