use std::task::{Context, Poll};
use glaredb_core::arrays::array::Array;
use glaredb_core::arrays::array::physical_type::{
AddressableMut,
MutableScalarStorage,
PhysicalBool,
PhysicalF64,
PhysicalI64,
PhysicalUtf8,
};
use glaredb_core::arrays::batch::Batch;
use glaredb_core::arrays::datatype::DataTypeId;
use glaredb_core::execution::operators::PollPull;
use glaredb_core::functions::cast::parse::{BoolParser, Float64Parser, Int64Parser, Parser};
use glaredb_core::runtime::filesystem::AnyFile;
use glaredb_core::runtime::filesystem::file_provider::MultiFileProvider;
use glaredb_core::storage::projections::{ProjectedColumn, Projections};
use glaredb_error::{DbError, Result, ResultExt};
use crate::decoder::{ByteRecords, CsvDecoder};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct CsvShape {
pub has_header: bool,
pub num_columns: usize,
}
#[derive(Debug)]
pub struct CsvReader {
file: Option<AnyFile>,
shape: CsvShape,
read_buf: Vec<u8>,
records: ByteRecords,
decoder: CsvDecoder,
projections: Projections,
state: ReaderState,
current_count: i64,
}
#[derive(Debug)]
enum ReaderState {
Reading {
skip_first: bool,
},
Flushing {
record_offset: usize,
stream_exhausted: bool,
},
Exhausted,
}
impl CsvReader {
pub fn new(
shape: CsvShape,
projections: Projections,
read_buf: Vec<u8>,
decoder: CsvDecoder,
records: ByteRecords,
) -> Self {
CsvReader {
file: None,
shape,
read_buf,
records,
decoder,
projections,
state: ReaderState::Reading {
skip_first: shape.has_header,
},
current_count: 0,
}
}
pub fn prepare(&mut self, file: AnyFile) {
self.file = Some(file);
self.records.clear_all(); self.state = ReaderState::Reading {
skip_first: self.shape.has_header,
};
self.current_count = 0;
}
pub fn poll_pull(&mut self, cx: &mut Context, output: &mut Batch) -> Result<PollPull> {
let out_cap = output.write_capacity()?;
debug_assert_ne!(0, out_cap);
let file = match self.file.as_mut() {
Some(file) => file,
None => {
return Err(DbError::new(
"Attempted to pull from CSV reader without preparing a file",
));
}
};
loop {
match self.state {
ReaderState::Reading { skip_first } => {
match file.call_poll_read(cx, &mut self.read_buf)? {
Poll::Ready(0) => {
self.state = ReaderState::Flushing {
record_offset: if skip_first { 1 } else { 0 },
stream_exhausted: true,
}
}
Poll::Ready(n) => {
let _ = self.decoder.decode(&self.read_buf[0..n], &mut self.records);
if self.records.num_records() >= out_cap {
self.state = ReaderState::Flushing {
record_offset: if skip_first { 1 } else { 0 },
stream_exhausted: false,
}
}
}
Poll::Pending => return Ok(PollPull::Pending),
}
}
ReaderState::Flushing {
record_offset,
stream_exhausted,
} => {
if record_offset >= self.records.num_records() {
if stream_exhausted {
self.state = ReaderState::Exhausted;
} else {
self.records.clear_completed();
debug_assert_eq!(0, self.records.num_records());
self.state = ReaderState::Reading {
skip_first: false,
};
}
continue;
}
let remaining = self.records.num_records() - record_offset;
let write_count = usize::min(remaining, out_cap);
self.write_batch(record_offset, output, 0, write_count)?;
self.state = ReaderState::Flushing {
record_offset: record_offset + write_count,
stream_exhausted,
};
self.current_count += write_count as i64;
output.set_num_rows(write_count)?;
return Ok(PollPull::HasMore);
}
ReaderState::Exhausted => {
output.set_num_rows(0)?;
return Ok(PollPull::Exhausted);
}
}
}
}
fn write_batch(
&self,
records_offset: usize,
batch: &mut Batch,
write_offset: usize,
count: usize,
) -> Result<()> {
self.projections
.for_each_column(batch, &mut |col_idx, array| match col_idx {
ProjectedColumn::Data(col_idx) => {
match array.datatype().id() {
DataTypeId::Boolean => self.write_primitive::<PhysicalBool, _>(
records_offset,
col_idx,
array,
write_offset,
count,
BoolParser,
)?,
DataTypeId::Int64 => self.write_primitive::<PhysicalI64, _>(
records_offset,
col_idx,
array,
write_offset,
count,
Int64Parser::new(),
)?,
DataTypeId::Float64 => self.write_primitive::<PhysicalF64, _>(
records_offset,
col_idx,
array,
write_offset,
count,
Float64Parser::new(),
)?,
DataTypeId::Utf8 => {
self.write_string(records_offset, col_idx, array, write_offset, count)?
}
other => {
return Err(DbError::new("Unhandled datatype for csv scanning")
.with_field("datatype", other));
}
}
Ok(())
}
ProjectedColumn::Metadata(MultiFileProvider::META_PROJECTION_FILENAME) => {
let file = self
.file
.as_ref()
.expect("file to be Some when writing projections");
let data = PhysicalUtf8::buffer_downcast_mut(array.data_mut())?;
let indices = write_offset..(write_offset + count);
data.put_duplicated(file.call_path().as_bytes(), indices)?;
Ok(())
}
ProjectedColumn::Metadata(MultiFileProvider::META_PROJECTION_ROWID) => {
let data = PhysicalI64::buffer_downcast_mut(array.data_mut())?;
let row_ids = &mut data.as_slice_mut()[write_offset..(write_offset + count)];
for (idx, row_id) in row_ids.iter_mut().enumerate() {
*row_id = self.current_count + idx as i64;
}
Ok(())
}
other => panic!("invalid projection: {other:?}"),
})?;
Ok(())
}
fn write_primitive<S, P>(
&self,
records_offset: usize,
field_idx: usize,
array: &mut Array,
write_offset: usize,
count: usize,
mut parser: P,
) -> Result<()>
where
S: MutableScalarStorage,
S::StorageType: Sized,
P: Parser<Type = S::StorageType>,
{
let (data, validity) = array.data_and_validity_mut();
let mut output = S::get_addressable_mut(data)?;
for idx in 0..count {
let record_idx = idx + records_offset;
let write_idx = idx + write_offset;
let record = self.records.get_record(record_idx);
if record.num_fields() != self.shape.num_columns {
return Err(DbError::new(format!(
"Expected {} columns in file, got {}",
self.shape.num_columns,
record.num_fields()
)));
}
let field = record.field(field_idx).unwrap();
let field = std::str::from_utf8(field).context("failed to parse field as utf8")?;
if field.is_empty() {
validity.set_invalid(write_idx);
} else {
let v = parser.parse(field).ok_or_else(|| {
DbError::new(format!("Failed to parse '{field}' as {}", S::PHYSICAL_TYPE)) })?;
output.put(write_idx, &v);
}
}
Ok(())
}
fn write_string(
&self,
records_offset: usize,
field_idx: usize,
array: &mut Array,
write_offset: usize,
count: usize,
) -> Result<()> {
let (data, validity) = array.data_and_validity_mut();
let mut output = PhysicalUtf8::get_addressable_mut(data)?;
for idx in 0..count {
let record_idx = idx + records_offset;
let write_idx = idx + write_offset;
let record = self.records.get_record(record_idx);
if record.num_fields() != self.shape.num_columns {
return Err(DbError::new(format!(
"Expected {} columns in file, got {}",
self.shape.num_columns,
record.num_fields()
)));
}
let field = record.field(field_idx).unwrap();
let field = std::str::from_utf8(field).context("failed to parse field as utf8")?;
if field.is_empty() {
validity.set_invalid(write_idx);
} else {
output.put(write_idx, field);
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use glaredb_core::arrays::datatype::DataType;
use glaredb_core::buffer::buffer_manager::DefaultBufferManager;
use glaredb_core::generate_batch;
use glaredb_core::runtime::filesystem::memory::MemoryFileHandle;
use glaredb_core::testutil::arrays::assert_batches_eq;
use glaredb_core::util::task::noop_context;
use super::*;
use crate::dialect::DialectOptions;
fn make_file(bytes: impl AsRef<[u8]>) -> AnyFile {
let file = MemoryFileHandle::from_bytes(&DefaultBufferManager, bytes).unwrap();
AnyFile::from_file(file)
}
#[test]
fn default_dialect_no_skip_header_complete_read() {
let input = r#"mario,9.5,8000
wario,10.0,950
yoshi,4.5,10000
"#;
let file = make_file(input);
let decoder = CsvDecoder::new(DialectOptions::default());
let output = ByteRecords::with_buffer_capacity(16);
let mut reader = CsvReader::new(
CsvShape {
has_header: false,
num_columns: 3,
},
Projections::new([0, 1, 2]),
vec![0; 256],
decoder,
output,
);
reader.prepare(file);
let mut batch = Batch::new(
[DataType::utf8(), DataType::float64(), DataType::int64()],
16,
)
.unwrap();
let poll = reader.poll_pull(&mut noop_context(), &mut batch).unwrap();
assert_eq!(PollPull::HasMore, poll);
let expected = generate_batch!(
["mario", "wario", "yoshi"],
[9.5, 10.0, 4.5],
[8000_i64, 950, 10000]
);
assert_batches_eq(&expected, &batch);
let poll = reader.poll_pull(&mut noop_context(), &mut batch).unwrap();
assert_eq!(PollPull::Exhausted, poll);
assert_eq!(0, batch.num_rows());
}
#[test]
fn default_dialect_no_skip_header_complete_read_small_read_buffer() {
let input = r#"mario,9.5,8000
wario,10.0,950
yoshi,4.5,10000
"#;
let file = make_file(input);
let decoder = CsvDecoder::new(DialectOptions::default());
let output = ByteRecords::with_buffer_capacity(16);
let mut reader = CsvReader::new(
CsvShape {
has_header: false,
num_columns: 3,
},
Projections::new([0, 1, 2]),
vec![0; 16],
decoder,
output,
);
reader.prepare(file);
let mut batch = Batch::new(
[DataType::utf8(), DataType::float64(), DataType::int64()],
16,
)
.unwrap();
let poll = reader.poll_pull(&mut noop_context(), &mut batch).unwrap();
assert_eq!(PollPull::HasMore, poll);
let expected = generate_batch!(
["mario", "wario", "yoshi"],
[9.5, 10.0, 4.5],
[8000_i64, 950, 10000]
);
assert_batches_eq(&expected, &batch);
let poll = reader.poll_pull(&mut noop_context(), &mut batch).unwrap();
assert_eq!(PollPull::Exhausted, poll);
assert_eq!(0, batch.num_rows());
}
#[test]
fn default_dialect_no_skip_header_large_read_buf_small_output_batch() {
let input = r#"mario,9.5,8000
wario,10.0,950
yoshi,4.5,10000
"#;
let file = make_file(input);
let decoder = CsvDecoder::new(DialectOptions::default());
let output = ByteRecords::with_buffer_capacity(16);
let mut reader = CsvReader::new(
CsvShape {
has_header: false,
num_columns: 3,
},
Projections::new([0, 1, 2]),
vec![0; 256],
decoder,
output,
);
reader.prepare(file);
let mut batch = Batch::new(
[DataType::utf8(), DataType::float64(), DataType::int64()],
2,
)
.unwrap();
let poll = reader.poll_pull(&mut noop_context(), &mut batch).unwrap();
assert_eq!(PollPull::HasMore, poll);
let expected = generate_batch!(["mario", "wario"], [9.5, 10.0], [8000_i64, 950]);
assert_batches_eq(&expected, &batch);
let poll = reader.poll_pull(&mut noop_context(), &mut batch).unwrap();
assert_eq!(PollPull::HasMore, poll);
println!("{}", batch.debug_table());
let expected = generate_batch!(["yoshi"], [4.5], [10000_i64]);
assert_batches_eq(&expected, &batch);
let poll = reader.poll_pull(&mut noop_context(), &mut batch).unwrap();
assert_eq!(PollPull::Exhausted, poll);
assert_eq!(0, batch.num_rows());
}
#[test]
fn default_dialect_no_skip_header_small_read_buf_small_output_batch() {
let input = r#"mario,9.5,8000
wario,10.0,950
yoshi,4.5,10000
"#;
let file = make_file(input);
let decoder = CsvDecoder::new(DialectOptions::default());
let output = ByteRecords::with_buffer_capacity(16);
let mut reader = CsvReader::new(
CsvShape {
has_header: false,
num_columns: 3,
},
Projections::new([0, 1, 2]),
vec![0; 16],
decoder,
output,
);
reader.prepare(file);
let mut batch = Batch::new(
[DataType::utf8(), DataType::float64(), DataType::int64()],
2,
)
.unwrap();
let poll = reader.poll_pull(&mut noop_context(), &mut batch).unwrap();
assert_eq!(PollPull::HasMore, poll);
let expected = generate_batch!(["mario", "wario"], [9.5, 10.0], [8000_i64, 950]);
assert_batches_eq(&expected, &batch);
let poll = reader.poll_pull(&mut noop_context(), &mut batch).unwrap();
assert_eq!(PollPull::HasMore, poll);
let expected = generate_batch!(["yoshi"], [4.5], [10000_i64]);
assert_batches_eq(&expected, &batch);
let poll = reader.poll_pull(&mut noop_context(), &mut batch).unwrap();
assert_eq!(PollPull::Exhausted, poll);
assert_eq!(0, batch.num_rows());
}
#[test]
fn default_dialect_skip_header_complete_read() {
let input = r#"string,float,int
mario,9.5,8000
wario,10.0,950
yoshi,4.5,10000
"#;
let file = make_file(input);
let decoder = CsvDecoder::new(DialectOptions::default());
let output = ByteRecords::with_buffer_capacity(16);
let mut reader = CsvReader::new(
CsvShape {
has_header: true,
num_columns: 3,
},
Projections::new([0, 1, 2]),
vec![0; 256],
decoder,
output,
);
reader.prepare(file);
let mut batch = Batch::new(
[DataType::utf8(), DataType::float64(), DataType::int64()],
16,
)
.unwrap();
let poll = reader.poll_pull(&mut noop_context(), &mut batch).unwrap();
assert_eq!(PollPull::HasMore, poll);
let expected = generate_batch!(
["mario", "wario", "yoshi"],
[9.5, 10.0, 4.5],
[8000_i64, 950, 10000]
);
assert_batches_eq(&expected, &batch);
let poll = reader.poll_pull(&mut noop_context(), &mut batch).unwrap();
assert_eq!(PollPull::Exhausted, poll);
assert_eq!(0, batch.num_rows());
}
}