use std::io::{BufWriter, Write};
use byteorder::{LittleEndian, WriteBytesExt};
use rand::{SeedableRng, rngs::SmallRng};
use super::FileHeader;
use crate::{
Policy, RNG_SEED, SequencingRecord,
error::{Result, WriteError},
};
pub fn write_flag<W: Write>(writer: &mut W, flag: u64) -> Result<()> {
writer.write_u64::<LittleEndian>(flag)?;
Ok(())
}
pub fn write_buffer<W: Write>(writer: &mut W, ebuf: &[u64]) -> Result<()> {
ebuf.iter()
.try_for_each(|&x| writer.write_u64::<LittleEndian>(x))?;
Ok(())
}
#[derive(Clone)]
pub struct Encoder {
header: FileHeader,
sbuffer: Vec<u64>, xbuffer: Vec<u64>,
s_ibuf: Vec<u8>, x_ibuf: Vec<u8>,
policy: Policy,
rng: SmallRng,
}
impl Encoder {
#[must_use]
pub fn new(header: FileHeader) -> Self {
Self::with_policy(header, Policy::default())
}
#[must_use]
pub fn with_policy(header: FileHeader, policy: Policy) -> Self {
Self {
header,
policy,
sbuffer: Vec::default(),
xbuffer: Vec::default(),
s_ibuf: Vec::default(),
x_ibuf: Vec::default(),
rng: SmallRng::seed_from_u64(RNG_SEED),
}
}
#[must_use]
pub fn is_paired(&self) -> bool {
self.header.is_paired()
}
pub fn encode_single(&mut self, primary: &[u8]) -> Result<Option<&[u64]>> {
if primary.len() != self.header.slen as usize {
return Err(WriteError::UnexpectedSequenceLength {
expected: self.header.slen,
got: primary.len(),
}
.into());
}
self.clear();
if self.header.bits.encode(primary, &mut self.sbuffer).is_err() {
self.clear();
if self
.policy
.handle(primary, &mut self.s_ibuf, &mut self.rng)?
{
self.header.bits.encode(&self.s_ibuf, &mut self.sbuffer)?;
} else {
return Ok(None);
}
}
Ok(Some(&self.sbuffer))
}
pub fn encode_paired(
&mut self,
primary: &[u8],
extended: &[u8],
) -> Result<Option<(&[u64], &[u64])>> {
if primary.len() != self.header.slen as usize {
return Err(WriteError::UnexpectedSequenceLength {
expected: self.header.slen,
got: primary.len(),
}
.into());
}
if extended.len() != self.header.xlen as usize {
return Err(WriteError::UnexpectedSequenceLength {
expected: self.header.xlen,
got: extended.len(),
}
.into());
}
self.clear();
if self.header.bits.encode(primary, &mut self.sbuffer).is_err()
|| self
.header
.bits
.encode(extended, &mut self.xbuffer)
.is_err()
{
self.clear();
if self
.policy
.handle(primary, &mut self.s_ibuf, &mut self.rng)?
&& self
.policy
.handle(extended, &mut self.x_ibuf, &mut self.rng)?
{
self.header.bits.encode(&self.s_ibuf, &mut self.sbuffer)?;
self.header.bits.encode(&self.x_ibuf, &mut self.xbuffer)?;
} else {
return Ok(None);
}
}
Ok(Some((&self.sbuffer, &self.xbuffer)))
}
pub fn clear(&mut self) {
self.sbuffer.clear();
self.xbuffer.clear();
self.s_ibuf.clear();
self.x_ibuf.clear();
}
}
#[derive(Default)]
pub struct WriterBuilder {
header: Option<FileHeader>,
policy: Option<Policy>,
headless: Option<bool>,
}
impl WriterBuilder {
#[must_use]
pub fn header(mut self, header: FileHeader) -> Self {
self.header = Some(header);
self
}
#[must_use]
pub fn policy(mut self, policy: Policy) -> Self {
self.policy = Some(policy);
self
}
#[must_use]
pub fn headless(mut self, headless: bool) -> Self {
self.headless = Some(headless);
self
}
pub fn build<W: Write>(self, inner: W) -> Result<Writer<W>> {
let Some(header) = self.header else {
return Err(WriteError::MissingHeader.into());
};
Writer::new(
inner,
header,
self.policy.unwrap_or_default(),
self.headless.unwrap_or(false),
)
}
}
#[derive(Clone)]
pub struct Writer<W: Write> {
inner: W,
encoder: Encoder,
headless: bool,
}
impl<W: Write> Writer<W> {
pub fn new(mut inner: W, header: FileHeader, policy: Policy, headless: bool) -> Result<Self> {
if !headless {
header.write_bytes(&mut inner)?;
}
Ok(Self {
inner,
encoder: Encoder::with_policy(header, policy),
headless,
})
}
pub fn is_paired(&self) -> bool {
self.encoder.is_paired()
}
pub fn header(&self) -> FileHeader {
self.encoder.header
}
pub fn policy(&self) -> Policy {
self.encoder.policy
}
#[deprecated]
pub fn write_record(&mut self, flag: Option<u64>, primary: &[u8]) -> Result<bool> {
let has_flag = self.encoder.header.flags;
if let Some(sbuffer) = self.encoder.encode_single(primary)? {
if has_flag {
write_flag(&mut self.inner, flag.unwrap_or(0))?;
}
write_buffer(&mut self.inner, sbuffer)?;
Ok(true)
} else {
Ok(false)
}
}
#[deprecated]
pub fn write_paired_record(
&mut self,
flag: Option<u64>,
primary: &[u8],
extended: &[u8],
) -> Result<bool> {
let has_flag = self.encoder.header.flags;
if let Some((sbuffer, xbuffer)) = self.encoder.encode_paired(primary, extended)? {
if has_flag {
write_flag(&mut self.inner, flag.unwrap_or(0))?;
}
write_buffer(&mut self.inner, sbuffer)?;
write_buffer(&mut self.inner, xbuffer)?;
Ok(true)
} else {
Ok(false)
}
}
pub fn push(&mut self, record: SequencingRecord) -> Result<bool> {
let has_flag = self.encoder.header.flags;
if has_flag {
write_flag(&mut self.inner, record.flag().unwrap_or(0))?;
}
if self.encoder.header.is_paired() && !record.is_paired() {
return Err(WriteError::ConfigurationMismatch {
attribute: "paired",
expected: self.encoder.header.is_paired(),
actual: record.is_paired(),
}
.into());
}
if self.encoder.header.is_paired() {
if let Some((sbuffer, xbuffer)) = self
.encoder
.encode_paired(record.s_seq, record.x_seq.unwrap_or_default())?
{
write_buffer(&mut self.inner, sbuffer)?;
write_buffer(&mut self.inner, xbuffer)?;
Ok(true)
} else {
Ok(false)
}
} else if let Some(buffer) = self.encoder.encode_single(record.s_seq)? {
write_buffer(&mut self.inner, buffer)?;
Ok(true)
} else {
Ok(false)
}
}
pub fn into_inner(self) -> W {
self.inner
}
pub fn by_ref(&mut self) -> &mut W {
&mut self.inner
}
pub fn flush(&mut self) -> Result<()> {
self.inner.flush()?;
Ok(())
}
pub fn new_encoder(&self) -> Encoder {
let mut encoder = self.encoder.clone();
encoder.clear();
encoder
}
pub fn is_headless(&self) -> bool {
self.headless
}
pub fn ingest(&mut self, other: &mut Writer<Vec<u8>>) -> Result<()> {
let other_inner = other.by_ref();
self.inner.write_all(other_inner)?;
other_inner.clear();
Ok(())
}
}
pub struct StreamWriter<W: Write> {
writer: Writer<BufWriter<W>>,
}
impl<W: Write> StreamWriter<W> {
pub fn new(inner: W, header: FileHeader, policy: Policy, headless: bool) -> Result<Self> {
Self::with_capacity(inner, 8192, header, policy, headless)
}
pub fn with_capacity(
inner: W,
capacity: usize,
header: FileHeader,
policy: Policy,
headless: bool,
) -> Result<Self> {
let buffered = BufWriter::with_capacity(capacity, inner);
let writer = Writer::new(buffered, header, policy, headless)?;
Ok(Self { writer })
}
#[deprecated(note = "use `push` method with SequencingRecord instead")]
pub fn write_record(&mut self, flag: Option<u64>, primary: &[u8]) -> Result<bool> {
#[allow(deprecated)]
self.writer.write_record(flag, primary)
}
#[deprecated(note = "use `push` method with SequencingRecord instead")]
pub fn write_paired_record(
&mut self,
flag: Option<u64>,
primary: &[u8],
extended: &[u8],
) -> Result<bool> {
#[allow(deprecated)]
self.writer.write_paired_record(flag, primary, extended)
}
pub fn push(&mut self, record: SequencingRecord) -> Result<bool> {
self.writer.push(record)
}
pub fn flush(&mut self) -> Result<()> {
self.writer.flush()
}
pub fn into_inner(self) -> Result<W> {
let bufw = self.writer.into_inner();
match bufw.into_inner() {
Ok(inner) => Ok(inner),
Err(e) => Err(std::io::Error::from(e).into()),
}
}
}
#[derive(Default)]
pub struct StreamWriterBuilder {
header: Option<FileHeader>,
policy: Option<Policy>,
headless: Option<bool>,
buffer_capacity: Option<usize>,
}
impl StreamWriterBuilder {
#[must_use]
pub fn header(mut self, header: FileHeader) -> Self {
self.header = Some(header);
self
}
#[must_use]
pub fn policy(mut self, policy: Policy) -> Self {
self.policy = Some(policy);
self
}
#[must_use]
pub fn headless(mut self, headless: bool) -> Self {
self.headless = Some(headless);
self
}
#[must_use]
pub fn buffer_capacity(mut self, capacity: usize) -> Self {
self.buffer_capacity = Some(capacity);
self
}
pub fn build<W: Write>(self, inner: W) -> Result<StreamWriter<W>> {
let Some(header) = self.header else {
return Err(WriteError::MissingHeader.into());
};
let capacity = self.buffer_capacity.unwrap_or(8192);
StreamWriter::with_capacity(
inner,
capacity,
header,
self.policy.unwrap_or_default(),
self.headless.unwrap_or(false),
)
}
}
#[cfg(test)]
mod testing {
use std::{fs::File, io::BufWriter};
use super::*;
use crate::SequencingRecordBuilder;
use crate::bq::{FileHeaderBuilder, SIZE_HEADER};
#[test]
fn test_headless() -> Result<()> {
let inner = Vec::new();
let mut writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(32).build()?)
.headless(true)
.build(inner)?;
assert!(writer.is_headless());
let inner = writer.by_ref();
assert!(inner.is_empty());
Ok(())
}
#[test]
fn test_not_headless() -> Result<()> {
let inner = Vec::new();
let mut writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(32).build()?)
.build(inner)?;
assert!(!writer.is_headless());
let inner = writer.by_ref();
assert_eq!(inner.len(), SIZE_HEADER);
Ok(())
}
#[test]
fn test_stdout() -> Result<()> {
let writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(32).build()?)
.build(std::io::stdout())?;
assert!(!writer.is_headless());
Ok(())
}
#[test]
fn test_to_path() -> Result<()> {
let path = "test_to_path.file";
let inner = File::create(path).map(BufWriter::new)?;
let mut writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(32).build()?)
.build(inner)?;
assert!(!writer.is_headless());
let inner = writer.by_ref();
inner.flush()?;
std::fs::remove_file(path)?;
Ok(())
}
#[test]
fn test_stream_writer() -> Result<()> {
let inner = Vec::new();
let writer = StreamWriterBuilder::default()
.header(FileHeaderBuilder::new().slen(32).build()?)
.buffer_capacity(16384)
.build(inner)?;
let inner = writer.into_inner()?;
assert_eq!(inner.len(), SIZE_HEADER);
Ok(())
}
#[test]
fn test_encoder_new() {
let header = FileHeaderBuilder::new().slen(8).build().unwrap();
let encoder = Encoder::new(header);
assert!(matches!(encoder.policy, Policy::IgnoreSequence));
}
#[test]
fn test_encoder_encode_single_wrong_length() {
let header = FileHeaderBuilder::new().slen(8).build().unwrap();
let mut encoder = Encoder::new(header);
let result = encoder.encode_single(b"ACGT");
assert!(result.is_err());
}
#[test]
fn test_encoder_encode_single_invalid_ignored() {
let header = FileHeaderBuilder::new().slen(8).build().unwrap();
let mut encoder = Encoder::with_policy(header, Policy::IgnoreSequence);
let result = encoder.encode_single(b"ACGTNNNN").unwrap();
assert!(result.is_none());
}
#[test]
fn test_encoder_encode_single_invalid_corrected() {
let header = FileHeaderBuilder::new().slen(8).build().unwrap();
let mut encoder = Encoder::with_policy(header, Policy::SetToA);
let result = encoder.encode_single(b"ACGTNNNN").unwrap();
assert!(result.is_some());
}
#[test]
fn test_encoder_encode_paired_wrong_primary_length() {
let header = FileHeaderBuilder::new().slen(8).xlen(8).build().unwrap();
let mut encoder = Encoder::new(header);
let result = encoder.encode_paired(b"ACGT", b"ACGTACGT");
assert!(result.is_err());
}
#[test]
fn test_encoder_encode_paired_wrong_extended_length() {
let header = FileHeaderBuilder::new().slen(8).xlen(8).build().unwrap();
let mut encoder = Encoder::new(header);
let result = encoder.encode_paired(b"ACGTACGT", b"ACGT");
assert!(result.is_err());
}
#[test]
fn test_encoder_encode_paired_invalid_ignored() {
let header = FileHeaderBuilder::new().slen(8).xlen(8).build().unwrap();
let mut encoder = Encoder::with_policy(header, Policy::IgnoreSequence);
let result = encoder.encode_paired(b"ACGTNNNN", b"ACGTACGT").unwrap();
assert!(result.is_none());
}
#[test]
fn test_encoder_encode_paired_invalid_corrected() {
let header = FileHeaderBuilder::new().slen(8).xlen(8).build().unwrap();
let mut encoder = Encoder::with_policy(header, Policy::SetToA);
let result = encoder.encode_paired(b"ACGTNNNN", b"NNNNACGT").unwrap();
assert!(result.is_some());
}
#[test]
fn test_writer_builder_missing_header() {
let result = WriterBuilder::default().build(Vec::new());
assert!(result.is_err());
}
#[test]
#[allow(deprecated)]
fn test_write_record_deprecated() -> Result<()> {
let mut writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(8).build()?)
.build(Vec::new())?;
let wrote = writer.write_record(None, b"ACGTACGT")?;
assert!(wrote);
Ok(())
}
#[test]
#[allow(deprecated)]
fn test_write_record_deprecated_skipped() -> Result<()> {
let mut writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(8).build()?)
.build(Vec::new())?;
let wrote = writer.write_record(None, b"NNNNNNNN")?;
assert!(!wrote);
Ok(())
}
#[test]
#[allow(deprecated)]
fn test_write_paired_record_deprecated() -> Result<()> {
let mut writer = WriterBuilder::default()
.header(
FileHeaderBuilder::new()
.slen(8)
.xlen(8)
.flags(true)
.build()?,
)
.build(Vec::new())?;
let wrote = writer.write_paired_record(Some(5), b"ACGTACGT", b"TTGGCCAA")?;
assert!(wrote);
Ok(())
}
#[test]
fn test_push_single() -> Result<()> {
let mut writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(8).build()?)
.build(Vec::new())?;
let record = SequencingRecordBuilder::default()
.s_seq(b"ACGTACGT")
.build()?;
assert!(writer.push(record)?);
Ok(())
}
#[test]
fn test_push_single_invalid_skipped() -> Result<()> {
let mut writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(8).build()?)
.build(Vec::new())?;
let record = SequencingRecordBuilder::default()
.s_seq(b"NNNNNNNN")
.build()?;
assert!(!writer.push(record)?);
Ok(())
}
#[test]
fn test_push_paired() -> Result<()> {
let mut writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(8).xlen(8).build()?)
.build(Vec::new())?;
let record = SequencingRecordBuilder::default()
.s_seq(b"ACGTACGT")
.x_seq(b"TTGGCCAA")
.build()?;
assert!(writer.push(record)?);
Ok(())
}
#[test]
fn test_push_paired_invalid_skipped() -> Result<()> {
let mut writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(8).xlen(8).build()?)
.build(Vec::new())?;
let record = SequencingRecordBuilder::default()
.s_seq(b"NNNNNNNN")
.x_seq(b"TTGGCCAA")
.build()?;
assert!(!writer.push(record)?);
Ok(())
}
#[test]
fn test_push_paired_mismatch() -> Result<()> {
let mut writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(8).xlen(8).build()?)
.build(Vec::new())?;
let record = SequencingRecordBuilder::default()
.s_seq(b"ACGTACGT")
.build()?;
let result = writer.push(record);
assert!(result.is_err());
Ok(())
}
#[test]
fn test_push_with_flag() -> Result<()> {
let mut writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(8).flags(true).build()?)
.build(Vec::new())?;
let record = SequencingRecordBuilder::default()
.s_seq(b"ACGTACGT")
.flag(99)
.build()?;
assert!(writer.push(record)?);
Ok(())
}
#[test]
fn test_new_encoder() -> Result<()> {
let writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(8).build()?)
.build(Vec::new())?;
let encoder = writer.new_encoder();
assert!(encoder.sbuffer.is_empty());
Ok(())
}
#[test]
fn test_writer_flush() -> Result<()> {
let mut writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(8).build()?)
.build(Vec::new())?;
writer.flush()?;
Ok(())
}
#[test]
fn test_writer_ingest() -> Result<()> {
let mut main_writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(8).build()?)
.build(Vec::new())?;
let header = main_writer.header();
let mut other_writer = WriterBuilder::default()
.header(header)
.headless(true)
.build(Vec::new())?;
let record = SequencingRecordBuilder::default()
.s_seq(b"ACGTACGT")
.build()?;
other_writer.push(record)?;
main_writer.ingest(&mut other_writer)?;
assert!(other_writer.by_ref().is_empty());
assert_eq!(main_writer.by_ref().len(), SIZE_HEADER + 8);
Ok(())
}
#[test]
fn test_writer_policy_accessor() -> Result<()> {
let writer = WriterBuilder::default()
.header(FileHeaderBuilder::new().slen(8).build()?)
.policy(Policy::SetToA)
.build(Vec::new())?;
assert!(matches!(writer.policy(), Policy::SetToA));
Ok(())
}
#[test]
fn test_stream_writer_new() -> Result<()> {
let writer = StreamWriter::new(
Vec::new(),
FileHeaderBuilder::new().slen(8).build()?,
Policy::default(),
false,
)?;
let inner = writer.into_inner()?;
assert_eq!(inner.len(), SIZE_HEADER);
Ok(())
}
#[test]
#[allow(deprecated)]
fn test_stream_writer_deprecated_methods() -> Result<()> {
let mut writer = StreamWriter::new(
Vec::new(),
FileHeaderBuilder::new().slen(8).xlen(8).build()?,
Policy::default(),
false,
)?;
assert!(writer.write_record(None, b"ACGTACGT")?);
assert!(writer.write_paired_record(None, b"ACGTACGT", b"TTGGCCAA")?);
writer.flush()?;
Ok(())
}
#[test]
fn test_stream_writer_push() -> Result<()> {
let mut writer = StreamWriter::new(
Vec::new(),
FileHeaderBuilder::new().slen(8).build()?,
Policy::default(),
false,
)?;
let record = SequencingRecordBuilder::default()
.s_seq(b"ACGTACGT")
.build()?;
assert!(writer.push(record)?);
writer.flush()?;
let inner = writer.into_inner()?;
assert_eq!(inner.len(), SIZE_HEADER + 8);
Ok(())
}
#[test]
fn test_stream_writer_builder_missing_header() {
let result = StreamWriterBuilder::default().build(Vec::new());
assert!(result.is_err());
}
#[test]
fn test_stream_writer_builder_with_policy_and_headless() -> Result<()> {
let inner = Vec::new();
let writer = StreamWriterBuilder::default()
.header(FileHeaderBuilder::new().slen(8).build()?)
.policy(Policy::SetToA)
.headless(true)
.build(inner)?;
let inner = writer.into_inner()?;
assert!(inner.is_empty());
Ok(())
}
}