mod decompression;
use crate::{
de::{
read::{Read, ReadSlice},
DeError,
},
object_container_file_encoding::{CompressionCodec, Metadata, HEADER_CONST, METADATA_SCHEMA},
*,
};
use {
decompression::DecompressionState,
serde::{
de::{DeserializeOwned, DeserializeSeed},
Deserialize,
},
std::{marker::PhantomData, sync::Arc},
};
pub struct Reader<R: de::read::take::Take> {
reader_state: ReaderState<'static, R>,
compression_codec: CompressionCodec,
sync_marker: [u8; 16],
pretend_eof_because_yielded_unrecoverable_error: bool,
schema: Arc<Schema>,
}
#[derive(Debug, thiserror::Error)]
pub enum FailedToInitializeReader {
#[error("Reader input is not an avro object container file: could not match the header")]
NotAvroObjectContainerFile,
#[error("Failed to validate avro object container file header: {}", _0)]
FailedToDeserializeHeader(DeError),
#[error("Failed to parse schema in avro object container file: {}", _0)]
FailedToParseSchema(schema::SchemaError),
}
impl<'a> Reader<de::read::SliceRead<'a>> {
pub fn from_slice(slice: &'a [u8]) -> Result<Self, FailedToInitializeReader> {
Self::new(de::read::SliceRead::new(slice))
}
}
impl<R: std::io::BufRead> Reader<de::read::ReaderRead<R>> {
pub fn from_reader(reader: R) -> Result<Self, FailedToInitializeReader> {
Self::new(de::read::ReaderRead::new(reader))
}
}
impl<R> Reader<R>
where
R: Read + de::read::take::Take + std::io::BufRead,
<R as de::read::take::Take>::Take: std::io::BufRead,
{
pub fn new<'de>(reader: R) -> Result<Self, FailedToInitializeReader>
where
R: ReadSlice<'de>,
{
Self::new_and_metadata::<()>(reader).map(|(reader, ())| reader)
}
pub fn new_and_metadata<'de, M>(mut reader: R) -> Result<(Self, M), FailedToInitializeReader>
where
R: ReadSlice<'de>,
M: Deserialize<'de>,
{
if reader
.read_const_size_buf::<4>()
.map_err(FailedToInitializeReader::FailedToDeserializeHeader)?
!= HEADER_CONST
{
return Err(FailedToInitializeReader::NotAvroObjectContainerFile);
}
let mut metadata_deserializer_config =
de::DeserializerConfig::from_schema_node(METADATA_SCHEMA);
metadata_deserializer_config.max_seq_size = 1_000;
let mut metadata_deserializer_state =
de::DeserializerState::with_config(reader, metadata_deserializer_config);
let metadata: Metadata<String, M> =
serde::Deserialize::deserialize(metadata_deserializer_state.deserializer())
.map_err(FailedToInitializeReader::FailedToDeserializeHeader)?;
reader = metadata_deserializer_state.into_reader();
let schema: Arc<Schema> = Arc::new(
metadata
.schema
.parse()
.map_err(FailedToInitializeReader::FailedToParseSchema)?,
);
let sync_marker = reader
.read_const_size_buf::<16>()
.map_err(FailedToInitializeReader::FailedToDeserializeHeader)?;
let schema_root = unsafe { schema.root_with_fake_static_lifetime() };
Ok((
Self {
reader_state: ReaderState::NotInBlock {
reader,
config: de::DeserializerConfig::from_schema_node(schema_root),
decompression_buffer: Vec::new(),
},
compression_codec: metadata.codec,
sync_marker,
pretend_eof_because_yielded_unrecoverable_error: false,
schema,
},
metadata.user_metadata,
))
}
pub fn deserialize<'r, 'rs, T: DeserializeOwned>(
&'r mut self,
) -> impl Iterator<Item = Result<T, DeError>> + 'r
where
R: ReadSlice<'rs>,
<R as de::read::take::Take>::Take: ReadSlice<'rs>,
{
self.deserialize_inner()
}
pub fn deserialize_borrowed<'r, 'de, T: Deserialize<'de>>(
&'r mut self,
) -> impl Iterator<Item = Result<T, DeError>> + 'r
where
R: ReadSlice<'de> + IsSliceRead,
<R as de::read::take::Take>::Take: ReadSlice<'de>,
{
Self::deserialize_inner::<T>(self)
}
fn deserialize_inner<'r, 'de, T: Deserialize<'de>>(
&'r mut self,
) -> impl Iterator<Item = Result<T, DeError>> + 'r
where
R: ReadSlice<'de>,
<R as de::read::take::Take>::Take: ReadSlice<'de>,
{
std::iter::from_fn(|| self.deserialize_seed_next(PhantomData::<T>).transpose())
}
pub fn deserialize_next<'a, T: DeserializeOwned>(&mut self) -> Result<Option<T>, DeError>
where
R: ReadSlice<'a>,
<R as de::read::take::Take>::Take: ReadSlice<'a>,
{
self.deserialize_seed_next(PhantomData::<T>)
}
pub fn deserialize_next_borrowed<'de, T: Deserialize<'de>>(
&mut self,
) -> Result<Option<T>, DeError>
where
R: ReadSlice<'de> + IsSliceRead,
<R as de::read::take::Take>::Take: ReadSlice<'de>,
{
self.deserialize_seed_next(PhantomData::<T>)
}
pub fn deserialize_seed_next<'de, S: DeserializeSeed<'de>>(
&mut self,
deserialize_seed: S,
) -> Result<Option<S::Value>, DeError>
where
R: ReadSlice<'de>,
<R as de::read::take::Take>::Take: ReadSlice<'de>,
{
if self.pretend_eof_because_yielded_unrecoverable_error {
return Ok(None);
}
let res = self.deserialize_next_inner(deserialize_seed);
if let Err(ref de_error) = res {
if de_error.io_error().is_some() || matches!(self.reader_state, ReaderState::Broken) {
self.pretend_eof_because_yielded_unrecoverable_error = true;
}
}
res
}
fn deserialize_next_inner<'de, S: DeserializeSeed<'de>>(
&mut self,
deserialize_seed: S,
) -> Result<Option<S::Value>, DeError>
where
R: ReadSlice<'de>,
<R as de::read::take::Take>::Take: ReadSlice<'de>,
{
loop {
match &mut self.reader_state {
ReaderState::Broken => {
return Err(DeError::new(
"Object container file reader is broken after error",
))
}
ReaderState::NotInBlock { reader, .. } => {
if reader
.fill_buf()
.map(|b| b.is_empty())
.map_err(DeError::io)?
{
break Ok(None);
}
let (mut reader, config, decompression_buffer) =
match std::mem::replace(&mut self.reader_state, ReaderState::Broken) {
ReaderState::NotInBlock {
reader,
config,
decompression_buffer,
} => (reader, config, decompression_buffer),
_ => unreachable!(),
};
let n_objects_in_block: i64 = reader.read_varint()?;
let n_objects_in_block: usize = n_objects_in_block
.try_into()
.map_err(|_| DeError::new("Invalid container file block object count"))?;
let block_size: i64 = reader.read_varint()?;
let block_size: usize = block_size
.try_into()
.map_err(|_| DeError::new("Invalid container file block size in bytes"))?;
let codec_data = self.compression_codec.state(
reader,
config,
decompression_buffer,
block_size,
)?;
self.reader_state = ReaderState::InBlock {
codec_data,
n_objects_in_block,
};
}
ReaderState::InBlock {
codec_data,
n_objects_in_block,
} => match n_objects_in_block.checked_sub(1) {
None => match std::mem::replace(&mut self.reader_state, ReaderState::Broken) {
ReaderState::InBlock {
codec_data,
n_objects_in_block: _,
} => {
let (mut reader, config, decompression_buffer) =
codec_data.into_source_reader_and_config()?;
let sync_marker = reader.read_const_size_buf::<16>()?;
if sync_marker != self.sync_marker {
return Err(DeError::new("Incorrect sync marker at end of block"));
}
self.reader_state = ReaderState::NotInBlock {
reader,
config,
decompression_buffer,
}
}
_ => unreachable!(),
},
Some(next_n_in_block) => {
*n_objects_in_block = next_n_in_block;
break match codec_data {
DecompressionState::Null {
deserializer_state, ..
} => deserialize_seed.deserialize(deserializer_state.deserializer()),
#[cfg(any(
feature = "deflate",
feature = "bzip2",
feature = "xz",
feature = "zstandard"
))]
DecompressionState::BufReader {
deserializer_state, ..
} => deserialize_seed.deserialize(deserializer_state.deserializer()),
#[cfg(feature = "snappy")]
DecompressionState::DecompressedOnConstruction {
deserializer_state,
..
} => deserialize_seed.deserialize(deserializer_state.deserializer()),
}
.map(Some);
}
},
}
}
}
pub fn schema(&self) -> &Arc<Schema> {
&self.schema
}
}
enum ReaderState<'s, R: de::read::take::Take> {
Broken,
NotInBlock {
reader: R,
config: de::DeserializerConfig<'s>,
decompression_buffer: Vec<u8>,
},
InBlock {
codec_data: DecompressionState<'s, R>,
n_objects_in_block: usize,
},
}
mod private {
pub trait IsSliceRead {}
}
use private::IsSliceRead;
impl IsSliceRead for de::read::SliceRead<'_> {}