use std::io::Write;
use std::sync::mpsc;
use std::sync::Arc;
use std::thread;
use std::time::Duration;
use anyhow::Result;
use parking_lot::Condvar;
use parking_lot::Mutex;
use crate::{
cli::{FilterOptions, OutputFormat},
output::{build_writers, Compression},
BUFFER_SIZE,
};
const DEFAULT_BUFFER_SIZE: usize = 1024 * 1024;
const MAXIMUM_BUFFER_SIZE: usize = 128 * 1024 * 1024;
const DEFAULT_SLEEP_MS: u64 = 100;
pub type BoxedWriter = Box<dyn Write + Send>;
pub type BoxedSegmentWriter = Box<dyn SegmentWriter + Send>;
pub trait SegmentWriter {
fn num_segments(&self) -> usize;
fn write_all_buffers(&mut self, buffers: &mut [Vec<u8>], counts: &mut [usize]) -> Result<()>;
fn generate_local_buffers(&self) -> Vec<Vec<u8>> {
vec![Vec::with_capacity(BUFFER_SIZE); self.num_segments()]
}
}
#[allow(clippy::too_many_arguments)]
pub fn build_segment_writer(
outdir: Option<&str>,
prefix: &str,
compression: Compression,
format: OutputFormat,
num_threads: usize,
filter_opts: &FilterOptions,
is_fifo: bool,
is_split: bool,
) -> Result<BoxedSegmentWriter> {
if is_split {
if is_fifo {
let wtr = BufferedWriter::new(
outdir,
prefix,
compression,
format,
num_threads,
filter_opts,
is_fifo,
)?;
Ok(Box::new(wtr))
} else {
let wtr = DirectWriter::new(
outdir,
prefix,
compression,
format,
num_threads,
filter_opts,
is_fifo,
)?;
Ok(Box::new(wtr))
}
} else {
let wtr = DirectWriter::new(
None,
prefix,
compression,
format,
num_threads,
filter_opts,
false,
)?;
Ok(Box::new(wtr))
}
}
struct ThreadWriter {
buffer_pair: Arc<(Mutex<Vec<u8>>, Condvar)>,
shutdown_sender: mpsc::Sender<()>,
join_handle: Option<thread::JoinHandle<Result<()>>>,
}
impl ThreadWriter {
fn new(mut handle: BoxedWriter) -> Self {
let buffer_pair = Arc::new((Mutex::new(Vec::new()), Condvar::new()));
let buffer_pair_clone = Arc::clone(&buffer_pair);
let (shutdown_sender, shutdown_receiver) = mpsc::channel();
let join_handle = thread::spawn(move || -> Result<()> {
let (buffer, cvar) = &*buffer_pair_clone;
loop {
let mut guard = buffer.lock();
while guard.is_empty() {
if shutdown_receiver.try_recv().is_ok() {
return Ok(()); }
cvar.wait(&mut guard);
if shutdown_receiver.try_recv().is_ok() {
return Ok(()); }
}
let mut data = std::mem::take(&mut *guard);
drop(guard);
handle.write_all(data.drain(..).as_slice())?;
handle.flush()?;
}
});
ThreadWriter {
buffer_pair,
shutdown_sender,
join_handle: Some(join_handle),
}
}
fn ingest(&self, data: &[u8]) {
let (buffer, cvar) = &*self.buffer_pair;
loop {
let mut guard = buffer.lock();
if guard.len() <= MAXIMUM_BUFFER_SIZE {
guard.extend_from_slice(data);
cvar.notify_one();
break;
} else {
thread::sleep(Duration::from_millis(DEFAULT_SLEEP_MS));
}
}
}
}
impl Drop for ThreadWriter {
fn drop(&mut self) {
self.shutdown_sender
.send(())
.expect("Error in sending signal");
let (_buffer, cvar) = &*self.buffer_pair;
cvar.notify_all();
if let Some(handle) = self.join_handle.take() {
handle
.join()
.expect("Error in joining thread")
.expect("Error within thread");
}
}
}
pub struct BufferedWriter {
segment_buffers: Vec<Vec<u8>>,
thread_writers: Vec<ThreadWriter>,
}
impl BufferedWriter {
pub fn new(
outdir: Option<&str>,
prefix: &str,
compression: Compression,
format: OutputFormat,
num_threads: usize,
filter_opts: &FilterOptions,
is_fifo: bool,
) -> Result<Self> {
let segment_handles = build_writers(
outdir,
prefix,
compression,
format,
num_threads,
filter_opts,
is_fifo,
)?;
let segment_buffers = vec![Vec::with_capacity(DEFAULT_BUFFER_SIZE); segment_handles.len()];
let thread_writers = segment_handles.into_iter().map(ThreadWriter::new).collect();
Ok(Self {
segment_buffers,
thread_writers,
})
}
fn write_to_handles(&mut self) -> Result<()> {
for (writer, buf) in self
.thread_writers
.iter()
.zip(self.segment_buffers.iter_mut())
{
if !buf.is_empty() {
writer.ingest(buf.drain(..).as_slice());
}
}
Ok(())
}
}
impl SegmentWriter for BufferedWriter {
fn num_segments(&self) -> usize {
self.thread_writers.len()
}
fn write_all_buffers(&mut self, buffers: &mut [Vec<u8>], counts: &mut [usize]) -> Result<()> {
for (shared_buf, (local_buf, local_count)) in self
.segment_buffers
.iter_mut()
.zip(buffers.iter_mut().zip(counts.iter_mut()))
{
if *local_count == 0 {
continue;
}
shared_buf.extend_from_slice(local_buf);
local_buf.clear();
*local_count = 0;
}
self.write_to_handles()
}
}
pub struct DirectWriter {
segment_handles: Vec<BoxedWriter>,
}
impl DirectWriter {
pub fn new(
outdir: Option<&str>,
prefix: &str,
compression: Compression,
format: OutputFormat,
num_threads: usize,
filter_opts: &FilterOptions,
is_fifo: bool,
) -> Result<Self> {
let segment_handles = build_writers(
outdir,
prefix,
compression,
format,
num_threads,
filter_opts,
is_fifo,
)?;
Ok(Self { segment_handles })
}
}
impl SegmentWriter for DirectWriter {
fn num_segments(&self) -> usize {
self.segment_handles.len()
}
fn write_all_buffers(&mut self, buffers: &mut [Vec<u8>], counts: &mut [usize]) -> Result<()> {
for (handle, (local_buf, local_count)) in self
.segment_handles
.iter_mut()
.zip(buffers.iter_mut().zip(counts.iter_mut()))
{
if *local_count == 0 {
continue;
}
handle.write_all(local_buf.drain(..).as_slice())?;
handle.flush()?;
*local_count = 0;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::{self, Write};
use std::sync::{Arc, Mutex};
struct TestWriter {
data: Arc<Mutex<Vec<u8>>>,
}
impl Write for TestWriter {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
let mut guard = self.data.lock().unwrap();
guard.extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
#[test]
fn direct_writer_write_all_buffers_happy_path_with_empty_segment() {
let data1 = Arc::new(Mutex::new(Vec::new()));
let data2 = Arc::new(Mutex::new(Vec::new()));
let writer1: Box<dyn Write + Send> = Box::new(TestWriter {
data: data1.clone(),
});
let writer2: Box<dyn Write + Send> = Box::new(TestWriter {
data: data2.clone(),
});
let mut dw = DirectWriter {
segment_handles: vec![writer1, writer2],
};
let mut buffers = vec![b"ACGT".to_vec(), Vec::new()];
let mut counts = vec![1, 0];
dw.write_all_buffers(&mut buffers, &mut counts).unwrap();
assert!(buffers[0].is_empty());
assert!(buffers[1].is_empty());
assert_eq!(counts, vec![0, 0]);
let written1 = data1.lock().unwrap().clone();
let written2 = data2.lock().unwrap().clone();
assert_eq!(written1, b"ACGT");
assert!(written2.is_empty());
}
}