use std::io::{self, BufRead, Read};
use std::num::NonZero;
use std::thread::{self, JoinHandle};
use crossbeam_channel::{bounded, Receiver, Sender};
use crate::reader::read_full;
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,
};
const MIN_BUFFERS: usize = 8;
type Decoded = io::Result<Buffer>;
type DecodedTx = oneshot::Sender<Decoded>;
type DecodedRx = oneshot::Receiver<Decoded>;
type InflateTx = Sender<(Buffer, DecodedTx)>;
type InflateRx = Receiver<(Buffer, DecodedTx)>;
type OrderTx = Sender<DecodedRx>;
type OrderRx = Receiver<DecodedRx>;
type RecycleTx = Sender<Buffer>;
type RecycleRx = Receiver<Buffer>;
#[derive(Default)]
struct Buffer {
raw: Vec<u8>,
raw_len: usize,
data: Vec<u8>,
data_len: usize,
pos: usize,
}
enum State<R> {
Running {
reader_handle: JoinHandle<R>,
inflater_handles: Vec<JoinHandle<()>>,
order_rx: OrderRx,
recycle_tx: RecycleTx,
},
Done,
}
pub struct MultithreadedReader<R>
where
R: Read + Send + 'static,
{
state: State<R>,
buffer: Buffer,
}
impl<R> MultithreadedReader<R>
where
R: Read + Send + 'static,
{
pub fn new(inner: R) -> Self {
let workers = thread::available_parallelism().map_or(1, NonZero::get);
Self::with_worker_count(NonZero::new(workers).unwrap_or(NonZero::<usize>::MIN), inner)
}
pub fn with_worker_count(worker_count: NonZero<usize>, inner: R) -> Self {
let capacity = worker_count.get().max(MIN_BUFFERS);
let (inflate_tx, inflate_rx) = bounded::<(Buffer, DecodedTx)>(capacity);
let (order_tx, order_rx) = bounded::<DecodedRx>(capacity);
let (recycle_tx, recycle_rx) = bounded::<Buffer>(capacity);
for _ in 0..capacity {
recycle_tx.send(Buffer::default()).expect("seeding the recycle pool cannot fail");
}
let reader_handle = spawn_reader(inner, inflate_tx, order_tx, recycle_rx);
let inflater_handles = spawn_inflaters(worker_count.get(), inflate_rx);
Self {
state: State::Running { reader_handle, inflater_handles, order_rx, recycle_tx },
buffer: Buffer::default(),
}
}
pub fn finish(&mut self) -> io::Result<R> {
match std::mem::replace(&mut self.state, State::Done) {
State::Running { reader_handle, mut inflater_handles, order_rx, recycle_tx } => {
drop(recycle_tx);
drop(order_rx);
for handle in inflater_handles.drain(..) {
handle.join().map_err(|_| thread_panicked("inflater"))?;
}
reader_handle.join().map_err(|_| thread_panicked("reader"))
}
State::Done => Err(io::Error::new(io::ErrorKind::Other, "reader already finished")),
}
}
fn next_block(&mut self) -> io::Result<bool> {
let State::Running { order_rx, recycle_tx, .. } = &self.state else {
return Ok(false);
};
loop {
let Ok(decoded_rx) = order_rx.recv() else {
return Ok(false);
};
let buffer = decoded_rx.recv().map_err(|_| {
io::Error::new(io::ErrorKind::Other, "bgzf worker thread stopped")
})??;
let spent = std::mem::replace(&mut self.buffer, buffer);
self.buffer.pos = 0;
recycle_tx.send(spent).ok();
if self.buffer.data_len > 0 {
return Ok(true);
}
}
}
}
impl MultithreadedReader<std::fs::File> {
pub fn from_path<P: AsRef<std::path::Path>>(path: P) -> io::Result<Self> {
std::fs::File::open(path).map(Self::new)
}
}
impl<R> Read for MultithreadedReader<R>
where
R: Read + Send + 'static,
{
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
let mut copied = 0;
while copied < buf.len() {
if self.buffer.pos >= self.buffer.data_len && !self.next_block()? {
break;
}
let available = &self.buffer.data[self.buffer.pos..self.buffer.data_len];
let n = available.len().min(buf.len() - copied);
buf[copied..copied + n].copy_from_slice(&available[..n]);
self.buffer.pos += n;
copied += n;
}
Ok(copied)
}
}
impl<R> BufRead for MultithreadedReader<R>
where
R: Read + Send + 'static,
{
fn fill_buf(&mut self) -> io::Result<&[u8]> {
if self.buffer.pos >= self.buffer.data_len {
self.next_block()?;
}
Ok(&self.buffer.data[self.buffer.pos..self.buffer.data_len])
}
fn consume(&mut self, amt: usize) {
self.buffer.pos = (self.buffer.pos + amt).min(self.buffer.data_len);
}
}
impl<R> Drop for MultithreadedReader<R>
where
R: Read + Send + 'static,
{
fn drop(&mut self) {
if matches!(self.state, State::Running { .. }) {
let _ = self.finish();
}
}
}
fn thread_panicked(which: &str) -> io::Error {
io::Error::new(io::ErrorKind::Other, format!("bgzf {which} thread panicked"))
}
fn to_io(e: BgzfError) -> io::Error {
io::Error::new(io::ErrorKind::Other, e)
}
fn spawn_reader<R>(
mut reader: R,
inflate_tx: InflateTx,
order_tx: OrderTx,
recycle_rx: RecycleRx,
) -> JoinHandle<R>
where
R: Read + Send + 'static,
{
thread::spawn(move || {
while let Ok(mut buffer) = recycle_rx.recv() {
match read_raw_block(&mut reader, &mut buffer) {
Ok(false) => break, Ok(true) => {
let (decoded_tx, decoded_rx) = oneshot::channel::<Decoded>();
if inflate_tx.send((buffer, decoded_tx)).is_err()
|| order_tx.send(decoded_rx).is_err()
{
break;
}
}
Err(e) => {
let (decoded_tx, decoded_rx) = oneshot::channel::<Decoded>();
decoded_tx.send(Err(e)).ok();
order_tx.send(decoded_rx).ok();
break;
}
}
}
reader
})
}
fn spawn_inflaters(worker_count: usize, inflate_rx: InflateRx) -> Vec<JoinHandle<()>> {
(0..worker_count)
.map(|_| {
let inflate_rx = inflate_rx.clone();
thread::spawn(move || {
let mut decompressor = Decompressor::new();
while let Ok((mut buffer, decoded_tx)) = inflate_rx.recv() {
let result = decode_block(
&buffer.raw[..buffer.raw_len],
&mut buffer.data,
&mut decompressor,
);
let message = match result {
Ok(data_len) => {
buffer.data_len = data_len;
buffer.pos = 0;
Ok(buffer)
}
Err(e) => Err(e),
};
decoded_tx.send(message).ok();
}
})
})
.collect()
}
#[inline]
fn grow_to(buf: &mut Vec<u8>, len: usize) {
if buf.len() < len {
buf.resize(len, 0);
}
}
fn read_raw_block<R: Read>(reader: &mut R, buffer: &mut Buffer) -> io::Result<bool> {
let mut header = [0u8; BGZF_HEADER_SIZE];
if !read_full(reader, &mut header)? {
return Ok(false);
}
check_header(&header).map_err(to_io)?;
let block_size = get_block_size(&header);
if block_size < BGZF_HEADER_SIZE + BGZF_FOOTER_SIZE {
return Err(to_io(BgzfError::InvalidHeader("block size smaller than header plus footer")));
}
grow_to(&mut buffer.raw, block_size);
buffer.raw[..BGZF_HEADER_SIZE].copy_from_slice(&header);
reader.read_exact(&mut buffer.raw[BGZF_HEADER_SIZE..block_size])?;
buffer.raw_len = block_size;
Ok(true)
}
fn decode_block(
raw: &[u8],
data: &mut Vec<u8>,
decompressor: &mut Decompressor,
) -> io::Result<usize> {
let payload = &raw[BGZF_HEADER_SIZE..];
let check = get_footer_values(payload);
let expected_len = check.amount as usize;
if let Some(len) = stored_block_len(payload) {
if payload.len() == DEFLATE_STORED_HEADER_SIZE + len + BGZF_FOOTER_SIZE {
if len != expected_len {
return Err(to_io(BgzfError::InvalidHeader(
"stored block length disagrees with footer",
)));
}
let src = &payload[DEFLATE_STORED_HEADER_SIZE..DEFLATE_STORED_HEADER_SIZE + len];
let found = crc32(src);
if found != check.sum {
return Err(to_io(BgzfError::InvalidChecksum { found, expected: check.sum }));
}
grow_to(data, len);
data[..len].copy_from_slice(src);
return Ok(len);
}
}
if expected_len > MAX_BGZF_BLOCK_SIZE {
return Err(to_io(BgzfError::UncompressedSizeExceeded {
found: expected_len,
max: MAX_BGZF_BLOCK_SIZE,
}));
}
grow_to(data, expected_len);
decompressor
.decompress(strip_footer(payload), &mut data[..expected_len], check, true)
.map_err(to_io)?;
Ok(expected_len)
}
#[cfg(test)]
mod tests {
use std::io::{Cursor, Read, Write};
use crate::{CompressionLevel, Reader, Writer};
use super::*;
fn owned(blob: &[u8]) -> Cursor<Vec<u8>> {
Cursor::new(blob.to_vec())
}
fn make_bgzf(data: &[u8], level: u8) -> Vec<u8> {
let mut out = vec![];
let mut writer = Writer::new(&mut out, CompressionLevel::new(level).unwrap());
writer.write_all(data).unwrap();
writer.finish().unwrap();
out
}
fn sample(len: usize) -> Vec<u8> {
(0..len as u32).map(|i| i.wrapping_mul(2_654_435_761).rotate_left(13) as u8).collect()
}
fn read_serial(blob: &[u8]) -> Vec<u8> {
let mut out = vec![];
Reader::new(blob).read_to_end(&mut out).unwrap();
out
}
fn read_mt(blob: &[u8], workers: usize) -> Vec<u8> {
let mut out = vec![];
MultithreadedReader::with_worker_count(NonZero::new(workers).unwrap(), owned(blob))
.read_to_end(&mut out)
.unwrap();
out
}
#[test]
fn matches_single_threaded_at_each_worker_count() {
let input = sample(300_000); for level in [0u8, 1, 6] {
let blob = make_bgzf(&input, level);
let serial = read_serial(&blob);
assert_eq!(serial, input, "sanity: serial reader round-trips at level {level}");
for workers in [1usize, 2, 4, 8] {
assert_eq!(
read_mt(&blob, workers),
input,
"mt reader diverged at level {level}, {workers} workers"
);
}
}
}
#[test]
fn single_worker_round_trips() {
let input = sample(200_000);
let blob = make_bgzf(&input, 6);
assert_eq!(read_mt(&blob, 1), input);
}
#[test]
fn store_only_multi_block_round_trips() {
let input = sample(250_000);
let blob = make_bgzf(&input, 0);
for workers in [1usize, 3] {
assert_eq!(read_mt(&blob, workers), input);
}
}
#[test]
fn empty_stream_reads_nothing() {
let blob = make_bgzf(b"", 6);
assert!(read_mt(&blob, 4).is_empty());
}
#[test]
fn tiny_reads_reassemble_stream() {
let input = sample(150_000);
let blob = make_bgzf(&input, 6);
let mut reader =
MultithreadedReader::with_worker_count(NonZero::new(4).unwrap(), owned(&blob));
let mut out = vec![];
let mut byte = [0u8; 1];
while reader.read(&mut byte).unwrap() == 1 {
out.push(byte[0]);
}
assert_eq!(out, input);
}
#[test]
fn truncated_block_errors() {
let input = sample(200_000);
let blob = make_bgzf(&input, 6);
let truncated = &blob[..blob.len() - crate::BGZF_EOF.len() - 20];
let mut reader =
MultithreadedReader::with_worker_count(NonZero::new(4).unwrap(), owned(truncated));
let mut out = vec![];
assert!(
reader.read_to_end(&mut out).is_err(),
"a truncated trailing block must error, not read as EOF"
);
}
#[test]
fn corrupt_header_errors() {
let input = sample(120_000);
let mut blob = make_bgzf(&input, 6);
blob[12] = b'X';
let mut reader =
MultithreadedReader::with_worker_count(NonZero::new(2).unwrap(), owned(&blob));
let mut out = vec![];
assert!(reader.read_to_end(&mut out).is_err());
}
#[test]
fn corrupt_payload_errors() {
let input = sample(80_000);
let mut blob = make_bgzf(&input, 6);
blob[BGZF_HEADER_SIZE + 2] ^= 0xff;
let mut reader =
MultithreadedReader::with_worker_count(NonZero::new(4).unwrap(), owned(&blob));
let mut out = vec![];
assert!(reader.read_to_end(&mut out).is_err());
}
#[test]
fn early_drop_is_clean() {
let input = sample(500_000); let blob = make_bgzf(&input, 6);
let mut reader =
MultithreadedReader::with_worker_count(NonZero::new(4).unwrap(), owned(&blob));
let mut small = [0u8; 64];
let _ = reader.read(&mut small).unwrap();
drop(reader); }
#[test]
fn finish_returns_inner() {
let input = sample(100_000);
let blob = make_bgzf(&input, 6);
let mut reader = MultithreadedReader::new(owned(&blob));
let mut out = vec![];
reader.read_to_end(&mut out).unwrap();
assert_eq!(out, input);
reader.finish().expect("finish should join cleanly");
}
use proptest::prelude::*;
proptest! {
#![proptest_config(ProptestConfig::with_cases(48))]
#[test]
fn proptest_mt_reader_matches_serial(
input in prop::collection::vec(any::<u8>(), 1..100_000usize),
comp_level in 0..=12u8,
workers in 1usize..=4,
) {
let blob = make_bgzf(&input, comp_level);
let mut serial = vec![];
Reader::new(blob.as_slice()).read_to_end(&mut serial).unwrap();
prop_assert_eq!(&serial, &input);
let mut mt = vec![];
MultithreadedReader::with_worker_count(NonZero::new(workers).unwrap(), owned(&blob))
.read_to_end(&mut mt)
.unwrap();
prop_assert_eq!(mt, input);
}
}
}