use crate::error::{StreamError, StreamErrorKind};
use bytes::{Buf, BytesMut};
use csv_core::{ReadRecordResult, Reader as CoreReader};
use tokio_util::codec::Decoder;
const INITIAL_FIELDS: usize = 512;
const INITIAL_ENDS: usize = 16;
#[derive(Debug, Clone, Copy)]
pub struct CsvFrameConfig {
pub delimiter: u8,
pub quote: u8,
pub double_quote: bool,
pub escape: Option<u8>,
pub terminator: csv_core::Terminator,
}
impl CsvFrameConfig {
fn build(&self) -> CoreReader {
csv_core::ReaderBuilder::new()
.delimiter(self.delimiter)
.quote(self.quote)
.double_quote(self.double_quote)
.escape(self.escape)
.terminator(self.terminator)
.build()
}
}
#[derive(Debug)]
pub struct CsvRecordCodec {
core: CoreReader,
output: Vec<u8>,
ends: Vec<usize>,
outlen: usize,
endlen: usize,
header_pending: bool,
max_len: usize,
}
impl CsvRecordCodec {
pub fn new(config: CsvFrameConfig, has_headers: bool, max_len: usize) -> Self {
Self {
core: config.build(),
output: vec![0; INITIAL_FIELDS],
ends: vec![0; INITIAL_ENDS],
outlen: 0,
endlen: 0,
header_pending: has_headers,
max_len,
}
}
fn take_record(&mut self) -> csv::ByteRecord {
let mut fields: Vec<&[u8]> = Vec::with_capacity(self.endlen);
let mut start = 0;
for &end in &self.ends[..self.endlen] {
fields.push(&self.output[start..end]);
start = end;
}
let record = csv::ByteRecord::from(fields);
self.outlen = 0;
self.endlen = 0;
record
}
fn next_record(
&mut self,
buf: &mut BytesMut,
at_eof: bool,
) -> Result<Option<csv::ByteRecord>, StreamError> {
loop {
let input_was_empty = buf.is_empty();
if input_was_empty && !at_eof {
return Ok(None);
}
let (res, nin, nout, nend) = self.core.read_record(
&buf[..],
&mut self.output[self.outlen..],
&mut self.ends[self.endlen..],
);
buf.advance(nin);
self.outlen += nout;
self.endlen += nend;
let record_bytes = self
.outlen
.saturating_add(self.endlen.saturating_mul(std::mem::size_of::<usize>()));
if record_bytes > self.max_len {
return Err(StreamError::new(
StreamErrorKind::MaxLenReachedError,
None,
Some("Max record length reached".into()),
));
}
match res {
ReadRecordResult::InputEmpty => {
if input_was_empty {
return Ok(None);
}
continue;
}
ReadRecordResult::OutputFull => {
self.output.resize(self.output.len().saturating_mul(2).max(1), 0);
continue;
}
ReadRecordResult::OutputEndsFull => {
self.ends.resize(self.ends.len().saturating_mul(2).max(1), 0);
continue;
}
ReadRecordResult::Record => {
let record = self.take_record();
if self.header_pending {
self.header_pending = false;
continue;
}
return Ok(Some(record));
}
ReadRecordResult::End => return Ok(None),
}
}
}
}
impl Decoder for CsvRecordCodec {
type Item = csv::ByteRecord;
type Error = StreamError;
fn decode(&mut self, buf: &mut BytesMut) -> Result<Option<Self::Item>, StreamError> {
self.next_record(buf, false)
}
fn decode_eof(&mut self, buf: &mut BytesMut) -> Result<Option<Self::Item>, StreamError> {
self.next_record(buf, true)
}
}