use std::{
io::{self, Write},
thread::{JoinHandle, Scope, ScopedJoinHandle},
};
use bytes::{Bytes, BytesMut};
pub use flate2::Compression;
use flume::{bounded, Receiver, Sender};
use log::warn;
use crate::check::Check;
use crate::{CompressResult, FormatSpec, GzpError, Message, ZWriter, DICT_SIZE};
#[derive(Debug)]
pub struct ParCompressBuilder<F>
where
F: FormatSpec,
{
buffer_size: usize,
num_threads: usize,
compression_level: Compression,
format: F,
pin_threads: Option<usize>,
}
impl<F> ParCompressBuilder<F>
where
F: FormatSpec,
{
pub fn new() -> Self {
Self {
buffer_size: F::DEFAULT_BUFSIZE,
num_threads: num_cpus::get(),
compression_level: Compression::new(3),
format: F::new(),
pin_threads: None,
}
}
pub fn buffer_size(mut self, buffer_size: usize) -> Result<Self, GzpError> {
if buffer_size < DICT_SIZE {
return Err(GzpError::BufferSize(buffer_size, DICT_SIZE));
}
self.buffer_size = buffer_size;
Ok(self)
}
pub fn num_threads(mut self, num_threads: usize) -> Result<Self, GzpError> {
if num_threads == 0 {
return Err(GzpError::NumThreads(num_threads));
}
self.num_threads = num_threads;
Ok(self)
}
pub fn compression_level(mut self, compression_level: Compression) -> Self {
self.compression_level = compression_level;
self
}
pub fn pin_threads(mut self, pin_threads: Option<usize>) -> Self {
if core_affinity::get_core_ids().is_none() {
warn!("Pinning threads is not supported on your platform. Please see core_affinity_rs. No threads will be pinned, but everything will work.");
self.pin_threads = None;
} else {
self.pin_threads = pin_threads;
}
self
}
pub fn from_writer<W: Write + Send + 'static>(self, writer: W) -> ParCompress<'static, F, W> {
let (tx_compressor, rx_compressor) = bounded(self.num_threads * 2);
let (tx_writer, rx_writer) = bounded(self.num_threads * 2);
let buffer_size = self.buffer_size;
let comp_level = self.compression_level;
let pin_threads = self.pin_threads;
let format = self.format;
let num_threads = self.num_threads;
let handle = std::thread::spawn(move || {
ParCompress::run(
&rx_compressor,
&rx_writer,
writer,
num_threads,
comp_level,
format,
pin_threads,
)
});
ParCompress {
handle: Some(MaybeScopedJoinHandle::Static(handle)),
tx_compressor: Some(tx_compressor),
tx_writer: Some(tx_writer),
dictionary: None,
buffer: BytesMut::with_capacity(buffer_size),
buffer_size,
format,
}
}
pub fn from_borrowed_writer<'scope, 'env, W: Write + Send + 'scope>(
self,
writer: W,
scope: &'scope Scope<'scope, 'env>,
) -> ParCompress<'scope, F, W> {
let (tx_compressor, rx_compressor) = bounded(self.num_threads * 2);
let (tx_writer, rx_writer) = bounded(self.num_threads * 2);
let buffer_size = self.buffer_size;
let comp_level = self.compression_level;
let pin_threads = self.pin_threads;
let format = self.format;
let num_threads = self.num_threads;
let handle = scope.spawn(move || {
ParCompress::run(
&rx_compressor,
&rx_writer,
writer,
num_threads,
comp_level,
format,
pin_threads,
)
});
ParCompress {
handle: Some(MaybeScopedJoinHandle::Scoped(handle)),
tx_compressor: Some(tx_compressor),
tx_writer: Some(tx_writer),
dictionary: None,
buffer: BytesMut::with_capacity(buffer_size),
buffer_size,
format,
}
}
}
impl<F> Default for ParCompressBuilder<F>
where
F: FormatSpec,
{
fn default() -> Self {
Self::new()
}
}
enum MaybeScopedJoinHandle<'scope, T> {
Static(JoinHandle<T>),
Scoped(ScopedJoinHandle<'scope, T>),
}
impl<'scope, T> MaybeScopedJoinHandle<'scope, T> {
fn join(self) -> Result<T, Box<dyn std::any::Any + Send>> {
match self {
MaybeScopedJoinHandle::Static(handle) => handle.join(),
MaybeScopedJoinHandle::Scoped(handle) => handle.join(),
}
}
}
#[allow(unused)]
pub struct ParCompress<'scope, F, W>
where
F: FormatSpec,
W: Write,
{
handle: Option<MaybeScopedJoinHandle<'scope, Result<W, GzpError>>>,
tx_compressor: Option<Sender<Message<F::C>>>,
tx_writer: Option<Sender<Receiver<CompressResult<F::C>>>>,
buffer: BytesMut,
dictionary: Option<Bytes>,
buffer_size: usize,
format: F,
}
impl<'scope, F, W> ParCompress<'scope, F, W>
where
F: FormatSpec,
W: Write,
{
pub fn builder() -> ParCompressBuilder<F> {
ParCompressBuilder::new()
}
#[allow(clippy::needless_collect)]
fn run(
rx: &Receiver<Message<F::C>>,
rx_writer: &Receiver<Receiver<CompressResult<F::C>>>,
mut writer: W,
num_threads: usize,
compression_level: Compression,
format: F,
pin_threads: Option<usize>,
) -> Result<W, GzpError>
where
W: Write + Send,
{
let (core_ids, pin_threads) = if let Some(core_ids) = core_affinity::get_core_ids() {
(core_ids, pin_threads)
} else {
(vec![], None)
};
let handles: Vec<JoinHandle<Result<(), GzpError>>> = (0..num_threads)
.map(|i| {
let rx = rx.clone();
let core_ids = core_ids.clone();
std::thread::spawn(move || -> Result<(), GzpError> {
if let Some(pin_at) = pin_threads {
if let Some(id) = core_ids.get(pin_at + i) {
core_affinity::set_for_current(*id);
}
}
let mut compressor = format.create_compressor(compression_level)?;
while let Ok(m) = rx.recv() {
let chunk = &m.buffer;
let buffer = format.encode(
chunk,
&mut compressor,
compression_level,
m.dictionary.as_ref(),
m.is_last,
)?;
let mut check = F::create_check();
check.update(chunk);
m.oneshot
.send(Ok::<(F::C, Vec<u8>), GzpError>((check, buffer)))
.map_err(|_e| GzpError::ChannelSend)?;
}
Ok(())
})
})
.collect();
writer.write_all(&format.header(compression_level))?;
let mut running_check = F::create_check();
while let Ok(chunk_chan) = rx_writer.recv() {
let chunk_chan: Receiver<CompressResult<F::C>> = chunk_chan;
let (check, chunk) = chunk_chan.recv()??;
running_check.combine(&check);
writer.write_all(&chunk)?;
}
let footer = format.footer(&running_check);
writer.write_all(&footer)?;
writer.flush()?;
handles
.into_iter()
.try_for_each(|handle| match handle.join() {
Ok(result) => result,
Err(e) => std::panic::resume_unwind(e),
})?;
Ok(writer)
}
fn flush_last(&mut self, is_last: bool) -> std::io::Result<()> {
loop {
let b = self
.buffer
.split_to(std::cmp::min(self.buffer.len(), self.buffer_size))
.freeze();
let (mut m, r) = Message::new_parts(b, self.dictionary.take());
if is_last && self.buffer.is_empty() {
m.is_last = true;
}
if m.buffer.len() >= DICT_SIZE && !m.is_last && self.format.needs_dict() {
self.dictionary = Some(m.buffer.slice(m.buffer.len() - DICT_SIZE..));
}
let send_result = self.tx_writer.as_ref().unwrap().send(r);
if let Err(error) = send_result {
return Err(self.recover_send_error(error));
}
let send_result = self.tx_compressor.as_ref().unwrap().send(m);
if let Err(error) = send_result {
return Err(self.recover_send_error(error));
}
if self.buffer.is_empty() {
break;
}
}
Ok(())
}
#[cold]
fn recover_send_error<T>(&mut self, send_error: flume::SendError<T>) -> io::Error {
let handle = self.handle.take().unwrap();
drop(send_error);
drop(self.tx_compressor.take());
drop(self.tx_writer.take());
let error = match handle.join() {
Ok(result) => result.map(|_| ()),
Err(error) => std::panic::resume_unwind(error),
};
match error {
Ok(()) => std::panic::resume_unwind(Box::new(error)),
Err(GzpError::Io(error)) => error,
Err(error) => io::Error::other(error),
}
}
}
impl<'scope, F, W> ZWriter<W> for ParCompress<'scope, F, W>
where
F: FormatSpec,
W: Write,
{
fn finish(&mut self) -> Result<W, GzpError> {
self.flush_last(true)?;
drop(self.tx_compressor.take());
drop(self.tx_writer.take());
match self.handle.take().unwrap().join() {
Ok(result) => result,
Err(e) => std::panic::resume_unwind(e),
}
}
}
impl<'scope, F, W> Drop for ParCompress<'scope, F, W>
where
F: FormatSpec,
W: Write,
{
fn drop(&mut self) {
if self.tx_compressor.is_some() && self.tx_writer.is_some() && self.handle.is_some() {
self.finish().unwrap();
}
}
}
impl<'scope, F, W> Write for ParCompress<'scope, F, W>
where
F: FormatSpec,
W: Write,
{
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.buffer.extend_from_slice(buf);
while self.buffer.len() > self.buffer_size {
let b = self.buffer.split_to(self.buffer_size).freeze();
let (m, r) = Message::new_parts(b, self.dictionary.take());
self.dictionary = if self.format.needs_dict() {
Some(m.buffer.slice(m.buffer.len() - DICT_SIZE..))
} else {
None
};
let send_result = self.tx_writer.as_ref().unwrap().send(r);
if let Err(error) = send_result {
return Err(self.recover_send_error(error));
}
let send_result = self.tx_compressor.as_ref().unwrap().send(m);
if let Err(error) = send_result {
return Err(self.recover_send_error(error));
}
self.buffer
.reserve(self.buffer_size.saturating_sub(self.buffer.len()));
}
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
self.flush_last(false)
}
}
#[cfg(all(test, feature = "deflate"))]
mod tests {
use std::io::{self, Write};
use std::panic::{catch_unwind, AssertUnwindSafe};
use std::sync::mpsc;
use std::time::{Duration, Instant};
use crate::deflate::Gzip;
use crate::{GzpError, ZWriter};
use super::{MaybeScopedJoinHandle, ParCompress, ParCompressBuilder};
#[derive(Debug)]
struct FailingWriter {
write_attempted: mpsc::Sender<()>,
}
impl Write for FailingWriter {
fn write(&mut self, _buf: &[u8]) -> io::Result<usize> {
let _ = self.write_attempted.send(());
Err(io::Error::new(io::ErrorKind::BrokenPipe, "sink is gone"))
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
#[derive(Debug)]
struct PanickingWriter {
write_attempted: mpsc::Sender<()>,
}
impl Write for PanickingWriter {
fn write(&mut self, _buf: &[u8]) -> io::Result<usize> {
let _ = self.write_attempted.send(());
panic!("sink panicked");
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
fn wait_for_writer_thread<W: Write>(compressor: &ParCompress<'_, Gzip, W>) {
let deadline = Instant::now() + Duration::from_secs(5);
loop {
let is_finished = match compressor.handle.as_ref().unwrap() {
MaybeScopedJoinHandle::Static(handle) => handle.is_finished(),
MaybeScopedJoinHandle::Scoped(handle) => handle.is_finished(),
};
if is_finished {
return;
}
assert!(Instant::now() < deadline, "writer thread did not finish");
std::thread::yield_now();
}
}
fn compressor_with_failed_writer() -> ParCompress<'static, Gzip, FailingWriter> {
let (write_attempted, write_attempted_rx) = mpsc::channel();
let compressor = ParCompressBuilder::new()
.num_threads(1)
.unwrap()
.from_writer(FailingWriter { write_attempted });
write_attempted_rx
.recv_timeout(Duration::from_secs(5))
.expect("writer did not attempt the gzip header");
wait_for_writer_thread(&compressor);
compressor
}
fn assert_pipeline_closed<W: Write>(compressor: &ParCompress<'_, Gzip, W>) {
assert!(compressor.handle.is_none());
assert!(compressor.tx_compressor.is_none());
assert!(compressor.tx_writer.is_none());
}
fn assert_sink_error(error: GzpError) {
match error {
GzpError::Io(error) => {
assert_eq!(error.kind(), io::ErrorKind::BrokenPipe);
assert_eq!(error.to_string(), "sink is gone");
}
error => panic!("expected the sink's I/O error, got {:?}", error),
}
}
#[test]
fn finish_error_is_not_retried_by_drop() {
let mut compressor = compressor_with_failed_writer();
let result = catch_unwind(AssertUnwindSafe(move || {
let result = compressor.finish();
assert_pipeline_closed(&compressor);
drop(compressor);
result
}));
let error = result
.expect("dropping after a failed finish must not panic")
.expect_err("finish should report the sink error");
assert_sink_error(error);
}
#[test]
fn flush_error_is_not_retried_by_drop() {
let mut compressor = compressor_with_failed_writer();
let result = catch_unwind(AssertUnwindSafe(move || {
let result = compressor.flush();
assert_pipeline_closed(&compressor);
drop(compressor);
result
}));
let error = result
.expect("dropping after a failed flush must not panic")
.expect_err("flush should report the sink error");
assert_eq!(error.kind(), io::ErrorKind::BrokenPipe);
assert_eq!(error.to_string(), "sink is gone");
}
#[test]
fn write_error_still_recovers_the_sink_error() {
let mut compressor = compressor_with_failed_writer();
let input = vec![0; compressor.buffer_size + 1];
let result = catch_unwind(AssertUnwindSafe(move || {
let result = compressor.write(&input);
assert_pipeline_closed(&compressor);
drop(compressor);
result
}));
let error = result
.expect("dropping after a failed write must not panic")
.expect_err("write should report the sink error");
assert_eq!(error.kind(), io::ErrorKind::BrokenPipe);
assert_eq!(error.to_string(), "sink is gone");
}
#[test]
fn scoped_finish_error_is_not_retried_by_drop() {
let result = catch_unwind(AssertUnwindSafe(|| {
std::thread::scope(|scope| {
let (write_attempted, write_attempted_rx) = mpsc::channel();
let mut compressor = ParCompressBuilder::new()
.num_threads(1)
.unwrap()
.from_borrowed_writer(FailingWriter { write_attempted }, scope);
write_attempted_rx
.recv_timeout(Duration::from_secs(5))
.expect("writer did not attempt the gzip header");
wait_for_writer_thread(&compressor);
let error = compressor
.finish()
.expect_err("finish should report the sink error");
assert_pipeline_closed(&compressor);
assert_sink_error(error);
drop(compressor);
});
}));
result.expect("dropping a failed scoped compressor must not panic");
}
#[test]
fn writer_panic_is_resumed_only_once() {
let (write_attempted, write_attempted_rx) = mpsc::channel();
let mut compressor = ParCompressBuilder::new()
.num_threads(1)
.unwrap()
.from_writer(PanickingWriter { write_attempted });
write_attempted_rx
.recv_timeout(Duration::from_secs(5))
.expect("writer did not attempt the gzip header");
wait_for_writer_thread(&compressor);
let result = catch_unwind(AssertUnwindSafe(move || {
let panic = catch_unwind(AssertUnwindSafe(|| compressor.finish()));
assert_pipeline_closed(&compressor);
drop(compressor);
panic
}))
.expect("dropping after resuming the writer panic must not panic again");
let panic = result.expect_err("the writer panic should be resumed");
let message = panic
.downcast_ref::<&str>()
.copied()
.or_else(|| panic.downcast_ref::<String>().map(String::as_str));
assert_eq!(message, Some("sink panicked"));
}
}