use std::{
fs::File,
io::{self, Read},
path::Path,
};
use crate::{
check_header, crc32, get_block_size, get_footer_values, stored_block_len, strip_footer,
BgzfError, Decompressor, BGZF_FOOTER_SIZE, BGZF_HEADER_SIZE, DEFLATE_STORED_HEADER_SIZE,
MAX_BGZF_BLOCK_SIZE,
};
pub struct Reader<R>
where
R: Read,
{
decompressed_buffer: Vec<u8>,
compressed_buffer: Vec<u8>,
block_pos: usize,
block_end: usize,
header_buffer: Vec<u8>,
validate_crc: bool,
decompressor: Decompressor,
reader: R,
}
impl<R> Reader<R>
where
R: Read,
{
pub fn new(reader: R) -> Self {
Self {
decompressed_buffer: vec![0u8; MAX_BGZF_BLOCK_SIZE],
compressed_buffer: vec![0u8; MAX_BGZF_BLOCK_SIZE],
block_pos: 0,
block_end: 0,
header_buffer: vec![0u8; BGZF_HEADER_SIZE + DEFLATE_STORED_HEADER_SIZE],
validate_crc: true,
decompressor: Decompressor::new(),
reader,
}
}
#[must_use]
pub fn with_crc_validation(mut self, validate: bool) -> Self {
self.validate_crc = validate;
self
}
}
impl Reader<File> {
pub fn from_path<P>(path: P) -> io::Result<Self>
where
P: AsRef<Path>,
{
File::open(path).map(Self::new)
}
}
impl<R> Read for Reader<R>
where
R: Read,
{
#[inline]
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
let mut total_bytes_copied = 0;
loop {
let available = self.block_end - self.block_pos;
if available > 0 {
let n = available.min(buf.len() - total_bytes_copied);
buf[total_bytes_copied..total_bytes_copied + n]
.copy_from_slice(&self.decompressed_buffer[self.block_pos..self.block_pos + n]);
self.block_pos += n;
total_bytes_copied += n;
}
if total_bytes_copied == buf.len() {
break;
}
debug_assert!(total_bytes_copied < buf.len(), "More bytes copied than requested.");
if !read_full(&mut self.reader, &mut self.header_buffer)? {
break; }
check_header(&self.header_buffer[..BGZF_HEADER_SIZE])
.map_err(|e| io::Error::new(io::ErrorKind::Other, e))?;
let block_size = get_block_size(&self.header_buffer[..BGZF_HEADER_SIZE]);
if block_size < BGZF_HEADER_SIZE + BGZF_FOOTER_SIZE {
return Err(io::Error::new(
io::ErrorKind::Other,
BgzfError::InvalidHeader("block size smaller than header plus footer"),
));
}
let payload_len = block_size - BGZF_HEADER_SIZE;
let deflate_header: [u8; DEFLATE_STORED_HEADER_SIZE] = self.header_buffer
[BGZF_HEADER_SIZE..]
.try_into()
.expect("header_buffer holds the DEFLATE block header");
if let Some(len) = stored_block_len(&deflate_header) {
let with_footer = len + BGZF_FOOTER_SIZE;
if payload_len == DEFLATE_STORED_HEADER_SIZE + with_footer
&& with_footer <= self.decompressed_buffer.len()
{
self.reader.read_exact(&mut self.decompressed_buffer[..with_footer])?;
let check = get_footer_values(&self.decompressed_buffer[..with_footer]);
if check.amount as usize != len {
return Err(io::Error::new(
io::ErrorKind::Other,
BgzfError::InvalidHeader("stored block length disagrees with footer"),
));
}
if self.validate_crc {
let found = crc32(&self.decompressed_buffer[..len]);
if found != check.sum {
return Err(io::Error::new(
io::ErrorKind::Other,
BgzfError::InvalidChecksum { found, expected: check.sum },
));
}
}
self.block_pos = 0;
self.block_end = len;
continue;
}
}
self.compressed_buffer[..DEFLATE_STORED_HEADER_SIZE].copy_from_slice(&deflate_header);
self.reader
.read_exact(&mut self.compressed_buffer[DEFLATE_STORED_HEADER_SIZE..payload_len])?;
let compressed = &self.compressed_buffer[..payload_len];
let check = get_footer_values(compressed);
let decompressed_len = check.amount as usize;
if decompressed_len > self.decompressed_buffer.len() {
return Err(io::Error::new(
io::ErrorKind::Other,
BgzfError::UncompressedSizeExceeded {
found: decompressed_len,
max: self.decompressed_buffer.len(),
},
));
}
self.decompressor
.decompress(
strip_footer(compressed),
&mut self.decompressed_buffer[..decompressed_len],
check,
self.validate_crc,
)
.map_err(|e| io::Error::new(io::ErrorKind::Other, e))?;
self.block_pos = 0;
self.block_end = decompressed_len;
}
Ok(total_bytes_copied)
}
}
pub(crate) fn read_full<R: Read>(reader: &mut R, buf: &mut [u8]) -> io::Result<bool> {
let mut filled = 0;
while filled < buf.len() {
match reader.read(&mut buf[filled..]) {
Ok(0) if filled == 0 => return Ok(false),
Ok(0) => {
return Err(io::Error::new(io::ErrorKind::UnexpectedEof, "truncated BGZF block"))
}
Ok(n) => filled += n,
Err(ref e) if e.kind() == io::ErrorKind::Interrupted => {}
Err(e) => return Err(e),
}
}
Ok(true)
}