mod compression_codec;
use {
super::*,
de::{
read::{Read, ReadSlice},
DeError,
},
};
use {
compression_codec::{CompressionCodec, CompressionCodecState},
serde::{de::DeserializeOwned, Deserialize},
};
pub struct Reader<R: de::read::take::Take> {
reader_state: ReaderState<'static, R>,
compression_codec: CompressionCodec,
sync_marker: [u8; 16],
_schema: 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::ParseSchemaError),
}
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)?
!= [b'O', b'b', b'j', 1u8]
{
return Err(FailedToInitializeReader::NotAvroObjectContainerFile);
}
#[derive(serde_derive::Deserialize)]
struct Metadata<M> {
#[serde(rename = "avro.schema")]
schema: String,
#[serde(rename = "avro.codec")]
codec: CompressionCodec,
#[serde(flatten)]
user_metadata: M,
}
let mut metadata_deserializer_config = de::DeserializerConfig::from_schema_node(
&schema::SchemaNode::Map(&schema::SchemaNode::Bytes),
);
metadata_deserializer_config.max_seq_size = 1_000;
let mut metadata_deserializer_state =
de::DeserializerState::with_config(reader, metadata_deserializer_config);
let metadata: Metadata<M> =
serde::Deserialize::deserialize(metadata_deserializer_state.deserializer())
.map_err(FailedToInitializeReader::FailedToDeserializeHeader)?;
reader = metadata_deserializer_state.into_reader();
let schema: Schema = metadata
.schema
.parse()
.map_err(FailedToInitializeReader::FailedToParseSchema)?;
let sync_marker = reader
.read_const_size_buf::<16>()
.map_err(FailedToInitializeReader::FailedToDeserializeHeader)?;
let schema_root: &'static schema::SchemaNode<'static> = unsafe {
let schema = &*(&schema as *const Schema);
let a: *const schema::SchemaNode<'_> = schema.root() as *const schema::SchemaNode<'_>;
let b: *const schema::SchemaNode<'static> = a as *const _;
&*b
};
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,
_schema: 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_borrowed()
}
pub fn deserialize_borrowed<'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_next_borrowed().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_next_borrowed()
}
pub fn deserialize_next_borrowed<'de, T: Deserialize<'de>>(
&mut self,
) -> Result<Option<T>, 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!(),
}
return Ok(None);
}
Some(next_n_in_block) => {
*n_objects_in_block = next_n_in_block;
break match codec_data {
CompressionCodecState::Null {
deserializer_state, ..
} => T::deserialize(deserializer_state.deserializer()),
#[cfg(feature = "deflate")]
CompressionCodecState::Deflate {
deserializer_state, ..
} => T::deserialize(deserializer_state.deserializer()),
#[cfg(feature = "bzip2")]
CompressionCodecState::Bzip2 {
deserializer_state, ..
} => T::deserialize(deserializer_state.deserializer()),
#[cfg(feature = "snappy")]
CompressionCodecState::Snappy {
deserializer_state, ..
} => T::deserialize(deserializer_state.deserializer()),
#[cfg(feature = "xz")]
CompressionCodecState::Xz {
deserializer_state, ..
} => T::deserialize(deserializer_state.deserializer()),
#[cfg(feature = "zstandard")]
CompressionCodecState::Zstandard {
deserializer_state, ..
} => T::deserialize(deserializer_state.deserializer()),
}
.map(Some);
}
},
}
}
}
}
enum ReaderState<'s, R: de::read::take::Take> {
Broken,
NotInBlock {
reader: R,
config: de::DeserializerConfig<'s>,
decompression_buffer: Vec<u8>,
},
InBlock {
codec_data: CompressionCodecState<'s, R>,
n_objects_in_block: usize,
},
}