use std::io::{self, Write};
use std::path::Path;
use anyhow::{Context, Result};
use bgzf::{CompressionLevel, Writer as BgzfWriter};
use fgumi_raw_bam::RawRecord;
use noodles_sam::Header;
use noodles_sam::io::Writer as SamWriter;
use rawb_io::WriteBehind;
use tempfile::NamedTempFile;
const BAM_MAGIC: &[u8; 4] = b"BAM\x01";
enum Sink {
File(std::fs::File),
Stdout(io::Stdout),
Zstd(Box<zstd::stream::write::Encoder<'static, std::fs::File>>),
}
impl Sink {
fn finish(self) -> Result<()> {
match self {
Sink::File(mut f) => f.flush().context("flushing output file")?,
Sink::Stdout(mut s) => s.flush().context("flushing stdout")?,
Sink::Zstd(encoder) => {
encoder.finish().context("finishing zstd temp output")?;
}
}
Ok(())
}
}
impl Write for Sink {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
match self {
Sink::File(f) => f.write(buf),
Sink::Stdout(s) => s.write(buf),
Sink::Zstd(e) => e.write(buf),
}
}
fn flush(&mut self) -> io::Result<()> {
match self {
Sink::File(f) => f.flush(),
Sink::Stdout(s) => s.flush(),
Sink::Zstd(e) => e.flush(),
}
}
}
pub struct RawBamWriter {
bgzf: Option<BgzfWriter<WriteBehind<Sink>>>,
records: u64,
}
impl RawBamWriter {
pub fn open(
path: Option<&Path>,
header: &Header,
ring_bytes: usize,
level: CompressionLevel,
) -> Result<Self> {
let sink = match path {
Some(p) if p.to_string_lossy() != "-" => {
let f = std::fs::File::create(p)
.with_context(|| format!("creating {}", p.display()))?;
Sink::File(f)
}
_ => Sink::Stdout(io::stdout()),
};
Self::from_sink(sink, header, ring_bytes, level)
}
pub fn open_temp(
temp: &NamedTempFile,
header: &Header,
ring_bytes: usize,
zstd_level: Option<i32>,
) -> Result<Self> {
let file = temp.reopen().with_context(|| format!("reopening {}", temp.path().display()))?;
let sink = match zstd_level {
None => Sink::File(file),
Some(level) => {
let encoder = zstd::stream::write::Encoder::new(file, level)
.with_context(|| format!("starting zstd for {}", temp.path().display()))?;
Sink::Zstd(Box::new(encoder))
}
};
let stored = CompressionLevel::new(0).expect("compression level 0 is valid");
Self::from_sink(sink, header, ring_bytes, stored)
}
fn from_sink(
sink: Sink,
header: &Header,
ring_bytes: usize,
level: CompressionLevel,
) -> Result<Self> {
let threaded = WriteBehind::with_thread_name(sink, ring_bytes, "dupblaster");
let mut bgzf = BgzfWriter::new(threaded, level);
write_bam_header(&mut bgzf, header)?;
Ok(Self { bgzf: Some(bgzf), records: 0 })
}
pub fn write_record(&mut self, rec: &RawRecord) -> io::Result<()> {
let w = self.bgzf.as_mut().expect("writer already finished");
let block_size = u32::try_from(rec.len())
.map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "record exceeds u32"))?;
w.write_all(&block_size.to_le_bytes())?;
w.write_all(rec.as_ref())?;
self.records += 1;
Ok(())
}
pub fn records_written(&self) -> u64 {
self.records
}
pub fn finish(mut self) -> Result<()> {
if let Some(w) = self.bgzf.take() {
let threaded = w.finish().context("flushing BGZF writer")?;
let sink = threaded.finish().context("flushing IO writer thread")?;
sink.finish()?;
}
Ok(())
}
pub fn abandon(mut self) {
if let Some(mut w) = self.bgzf.take() {
let _ = w.flush();
std::mem::forget(w);
}
}
}
fn write_bam_header<W: Write>(writer: &mut W, header: &Header) -> Result<()> {
writer.write_all(BAM_MAGIC).context("writing BAM magic")?;
let text = serialize_sam_text(header)?;
let l_text =
u32::try_from(text.len()).map_err(|_| anyhow::anyhow!("BAM header text exceeds u32"))?;
writer.write_all(&l_text.to_le_bytes()).context("writing l_text")?;
writer.write_all(&text).context("writing SAM text")?;
let refs = header.reference_sequences();
let n_ref = u32::try_from(refs.len()).map_err(|_| anyhow::anyhow!("n_ref exceeds u32"))?;
writer.write_all(&n_ref.to_le_bytes()).context("writing n_ref")?;
for (name, map) in refs {
let l_name = u32::try_from(name.len() + 1)
.map_err(|_| anyhow::anyhow!("ref name length exceeds u32"))?;
writer.write_all(&l_name.to_le_bytes()).context("writing l_name")?;
writer.write_all(name).context("writing ref name")?;
writer.write_all(&[0]).context("writing ref name NUL")?;
let l_ref = u32::try_from(usize::from(map.length()))
.map_err(|_| anyhow::anyhow!("ref length exceeds u32"))?;
writer.write_all(&l_ref.to_le_bytes()).context("writing l_ref")?;
}
Ok(())
}
fn serialize_sam_text(header: &Header) -> Result<Vec<u8>> {
let mut buf = SamWriter::new(Vec::new());
buf.write_header(header).context("serializing SAM header text")?;
Ok(buf.into_inner())
}
#[cfg(test)]
mod tests {
use super::*;
fn temp_bam_bytes(level: Option<i32>, count: usize) -> u64 {
let temp = NamedTempFile::new().expect("temp file");
let mut writer = RawBamWriter::open_temp(&temp, &Header::default(), 1 << 20, level)
.expect("writer opens");
let mut record = RawRecord::new();
record.as_mut_vec().resize(64, 0);
for _ in 0..count {
writer.write_record(&record).expect("write");
}
writer.finish().expect("finish");
std::fs::metadata(temp.path()).expect("stat").len()
}
#[test]
fn compressing_the_temp_bam_makes_it_smaller() {
let plain = temp_bam_bytes(None, 20_000);
let compressed = temp_bam_bytes(Some(1), 20_000);
assert!(compressed < plain, "zstd grew the temp BAM: {compressed} vs {plain}");
}
#[test]
fn a_truncated_compressed_temp_bam_never_reads_back_complete() {
use crate::raw_reader::RawBamReader;
let temp = NamedTempFile::new().expect("temp file");
let written = 20_000u64;
let mut writer = RawBamWriter::open_temp(&temp, &Header::default(), 1 << 20, Some(1))
.expect("writer opens");
let mut record = RawRecord::new();
for i in 0..written {
let bytes = record.as_mut_vec();
bytes.clear();
bytes.extend((0..64u64).map(|b| (i.wrapping_mul(31).wrapping_add(b)) as u8));
writer.write_record(&record).expect("write");
}
assert_eq!(writer.records_written(), written);
writer.finish().expect("finish");
let full = std::fs::metadata(temp.path()).expect("stat").len();
std::fs::OpenOptions::new()
.write(true)
.open(temp.path())
.expect("reopen")
.set_len(full - full / 10)
.expect("truncate");
let file = std::fs::File::open(temp.path()).expect("open");
let decoder = zstd::stream::read::Decoder::new(file).expect("decoder");
let mut reader = RawBamReader::new(std::io::BufReader::new(decoder), false);
let mut recovered = 0u64;
let mut scratch = RawRecord::new();
let read_back = reader.read_header().map_err(|_| ()).and_then(|_| {
loop {
match reader.read_record(&mut scratch) {
Ok(true) => recovered += 1,
Ok(false) => return Ok(()),
Err(_) => return Err(()),
}
}
});
assert!(
read_back.is_err() || recovered < written,
"a truncated temp BAM reported success with all {written} records"
);
}
#[test]
fn the_writer_counts_the_records_it_wrote() {
let temp = NamedTempFile::new().expect("temp file");
let mut writer = RawBamWriter::open_temp(&temp, &Header::default(), 1 << 20, None)
.expect("writer opens");
let mut record = RawRecord::new();
record.as_mut_vec().resize(64, 0);
for _ in 0..7 {
writer.write_record(&record).expect("write");
}
assert_eq!(writer.records_written(), 7);
}
}