use std::{collections::VecDeque, time::Duration};
use anyhow::Context as _;
use create_records::{RecordMetadata, create_records, statistics::compute_statistics};
use rust_htslib::bcf;
use tokio::{
sync::{mpsc, watch::error::RecvError},
time::timeout,
};
use tracing::{Instrument, Span, debug, error, instrument, warn, warn_span};
use tracing_indicatif::span_ext::IndicatifSpanExt;
use generic_a_star::cost::AStarCost as _;
use crate::{
common::{
aligner::{
self,
result::{AlignmentFailure, SoftFailureReason, TwitcherAlignmentWithStatistics},
},
csv::{CSVAuxData, TwitcherCSVWriter, compute_record_hash},
},
counter,
vcf::pipeline::{
clusterizer::phasing::OutputPhasing,
record::InputRecord,
writer::{output_record::OutputRecord, region_writer::RegionWriter},
},
};
use super::Message;
mod create_records;
mod output_record;
pub mod region_writer;
pub struct Writer {
input: mpsc::Receiver<Message>,
output: SortingWriter,
region_output: Option<RegionWriter>,
csv_output: Option<TwitcherCSVWriter>,
}
impl Writer {
pub(super) const fn new(
input: mpsc::Receiver<Message>,
output: OutputWriter,
region_output: Option<RegionWriter>,
csv_output: Option<TwitcherCSVWriter>,
) -> Self {
Self {
input,
output: SortingWriter::new(1000, output),
region_output,
csv_output,
}
}
#[instrument(skip_all, fields(indicatif.pb_show = true))]
pub async fn run(mut self) -> anyhow::Result<()> {
let mut write_error: Option<anyhow::Error> = None;
loop {
let r = timeout(Duration::from_millis(100), self.input.recv()).await;
match r {
Ok(Some(m)) => match self.handle_message(m).await {
Ok(()) => {}
Err(e) => {
error!("Error in writer: {e}");
write_error.get_or_insert(e);
}
},
Ok(None) => break, Err(_) => {} }
Self::tick_progress();
}
self.output.flush_while(|_| true)?;
if let Some(mut rw) = self.region_output {
rw.flush().await?;
}
if let Some(e) = write_error {
return Err(e);
}
Ok(())
}
async fn handle_message(&mut self, m: Message) -> anyhow::Result<()> {
let (mut rx, aux, records, phasing) = match m {
Message::Passthrough(records) => {
return self.write_unchanged_records(records).map_err(Into::into);
}
Message::Cluster {
pending: None,
records,
phasing,
} => return self.reconstruct(records, phasing),
Message::Cluster {
pending: Some(pending),
records,
phasing,
} => {
let (rx, aux) = *pending;
(rx, aux, records, phasing)
}
};
let ts_alignment: Result<_, RecvError> =
async move {
loop {
let result = match timeout(Duration::from_millis(100), rx.recv()).await {
Ok(res) => res?,
Err(_elapsed) => {
Self::tick_progress();
continue;
}
};
counter!("alignments.finished").inc(1);
match &result.outcome {
Ok(realignment) if realignment.has_ts() => {
break Ok(Some(realignment.clone()));
}
Ok(_) => break Ok(None),
Err(AlignmentFailure::SoftFailure {
reason:
reason
@ (SoftFailureReason::OutOfMemory | SoftFailureReason::Timeout(_)),
}) => {
debug!(
"Alignment failed ({reason:?}); emitting the cluster without a template switch"
);
break Ok(None);
}
Err(AlignmentFailure::SoftFailure {
reason: SoftFailureReason::Other(error),
}) => {
warn!("Alignment failed: {error}");
break Ok(None);
}
Err(AlignmentFailure::Error { error }) => {
error!("Alignment failed: {error}");
break Ok(None);
}
}
}
}
.instrument(warn_span!("Handling Realignment", pos = %aux.region))
.await;
if let Some(realignment) = ts_alignment? {
self.write_one_alignment(realignment, aux, phasing, &records)
.await?;
} else {
self.reconstruct(records, phasing)?;
}
Ok(())
}
fn reconstruct(
&mut self,
records: Vec<(InputRecord, u32)>,
phasing: OutputPhasing,
) -> anyhow::Result<()> {
for (rec, allele_idx) in records {
let phasing_changed = phasing.changes(&rec, allele_idx);
let mut rec = self.output.adopt(rec);
{
let alleles = rec.alleles();
let reference = alleles
.first()
.context("Record without a reference allele")?
.to_vec();
let Some(alt) = alleles.get(allele_idx as usize).map(|a| a.to_vec()) else {
warn!(
"Reconstruct: record at {} has invalid allele index {allele_idx}; \
emitting unchanged",
rec.pos() + 1
);
self.output.write(rec, RecordProperties::OldRecord)?;
continue;
};
rec.set_alleles(&[&reference, &alt])?;
}
compute_statistics(&mut rec, None, Some(phasing), phasing_changed)?;
self.output.write(rec, RecordProperties::OldRecord)?;
}
Ok(())
}
async fn write_one_alignment(
&mut self,
realignment: TwitcherAlignmentWithStatistics,
aux: CSVAuxData,
phasing: OutputPhasing,
old_records: &[(InputRecord, u32)],
) -> Result<(), anyhow::Error> {
let TwitcherAlignmentWithStatistics { alignment, stats } = &realignment;
let variant_id = format!("{:032x}", compute_record_hash(&realignment, &aux));
let mut new_records = create_records(
alignment.alignment.iter_compact_cloned(),
(stats.reference_offset(), &aux.sequences.reference),
(stats.query_offset(), &aux.sequences.query),
&self.output.empty_record(),
aux.ref_context_region.start(),
RecordMetadata {
old_records: None,
alignment_cost: Some(alignment.cost.as_primitive()),
variant_id: Some(&variant_id),
cluster_grp: Some(aux.cluster_grp.as_str()),
phasing: Some(phasing),
phasing_changed: old_records
.iter()
.any(|(rec, allele_idx)| phasing.changes(rec, *allele_idx)),
},
)?;
if let Some(rw) = &mut self.region_output {
rw.write(new_records.iter().map(|(r, _)| r)).await?;
}
if let Some(csv) = &mut self.csv_output {
csv.write(&realignment, aux)?;
}
new_records.sort_by_key(|(rec, _)| rec.pos());
for (rec, prop) in new_records {
self.output.write(rec, prop)?;
}
Ok(())
}
fn tick_progress() {
let span = Span::current();
let total = counter!("alignments").get();
let progress = counter!("alignments.finished").get();
span.pb_set_length(total as u64);
span.pb_set_position(progress as u64);
span.pb_set_message(&format!(
"({} alignments running)",
aligner::RUNNING.load(std::sync::atomic::Ordering::Relaxed)
));
}
fn write_unchanged_records(
&mut self,
records: Vec<InputRecord>,
) -> Result<(), rust_htslib::errors::Error> {
for r in records {
let r = self.output.adopt(r);
self.output.write(r, RecordProperties::OldRecord)?;
}
Ok(())
}
}
pub struct SortingWriter {
buf: VecDeque<(OutputRecord, RecordProperties)>,
buf_coord_len: i64,
inner: OutputWriter,
}
impl SortingWriter {
const fn new(buf_coord_len: i64, inner: OutputWriter) -> Self {
Self {
buf: VecDeque::new(),
buf_coord_len,
inner,
}
}
fn adopt(&mut self, record: InputRecord) -> OutputRecord {
self.inner.adopt(record)
}
fn empty_record(&self) -> OutputRecord {
self.inner.empty_record()
}
fn write(
&mut self,
record: OutputRecord,
properties: RecordProperties,
) -> Result<(), rust_htslib::errors::Error> {
self.buf.push_back((record, properties));
self.swim_last();
self.flush()
}
fn swim_last(&mut self) {
let mut ix = self.buf.len().saturating_sub(1);
while ix >= 1 {
let (Some(this), Some(prev)) = (self.buf.get(ix), self.buf.get(ix - 1)) else {
return;
};
if (this.0.rid(), this.0.pos()) >= (prev.0.rid(), prev.0.pos()) {
return;
}
self.buf.swap(ix, ix - 1);
ix -= 1;
}
}
fn flush(&mut self) -> Result<(), rust_htslib::tpool::Error> {
let last = self
.buf
.back()
.and_then(|(rec, _)| Some((rec.rid()?, rec.pos())));
let len = self.buf_coord_len;
self.flush_while(|(rec, _)| {
last.is_some_and(|(last_rid, last_pos)| {
rec.rid()
.is_none_or(|rid| (last_rid, last_pos - len) > (rid, rec.pos()))
})
})
}
fn flush_while(
&mut self,
condition: impl Fn(&(OutputRecord, RecordProperties)) -> bool,
) -> Result<(), rust_htslib::tpool::Error> {
while self.buf.front().is_some_and(&condition) {
let Some((record, properties)) = self.buf.pop_front() else {
break;
};
tokio::task::block_in_place(|| self.inner.write(record, properties))?;
}
Ok(())
}
}
impl Drop for SortingWriter {
fn drop(&mut self) {
if let Err(err) = self.flush_while(|_| true) {
error!("Could not flush all the records: {err}");
}
}
}
enum RecordProperties {
OldRecord,
Realigned {
#[allow(unused)]
has_ts: bool,
},
}
pub enum OutputWriter {
Native { inner: bcf::Writer, only_ts: bool },
Buffered(FilteredOutputWriter),
}
impl OutputWriter {
pub const fn new_native(inner: bcf::Writer, only_ts: bool) -> Self {
Self::Native { inner, only_ts }
}
pub const fn new_buffered(inner: bcf::Writer, plus_minus: i64) -> Self {
Self::Buffered(FilteredOutputWriter {
buf: VecDeque::new(),
out: inner,
last_keep_pos: None,
current_rid: None,
current_pos: 0,
max_distance: plus_minus,
})
}
fn empty_record(&self) -> OutputRecord {
OutputRecord::new(self.inner().empty_record())
}
fn adopt(&mut self, record: InputRecord) -> OutputRecord {
let mut record = record.into_record();
self.inner_mut().translate(&mut record);
OutputRecord::new(record)
}
const fn inner(&self) -> &bcf::Writer {
match self {
Self::Native { inner, .. } => inner,
Self::Buffered(filtered_output_writer) => &filtered_output_writer.out,
}
}
const fn inner_mut(&mut self) -> &mut bcf::Writer {
match self {
Self::Native { inner, .. } => inner,
Self::Buffered(filtered_output_writer) => &mut filtered_output_writer.out,
}
}
fn write(
&mut self,
record: OutputRecord,
properties: RecordProperties,
) -> Result<(), rust_htslib::errors::Error> {
match (self, properties) {
(Self::Native { only_ts: true, .. }, RecordProperties::OldRecord) => {}
(
Self::Native {
inner,
only_ts: false,
},
RecordProperties::OldRecord,
)
| (Self::Native { inner, .. }, RecordProperties::Realigned { .. }) => {
inner.write(&record)?;
}
(Self::Buffered(filtered_output_writer), RecordProperties::OldRecord) => {
filtered_output_writer.write(record, false)?;
}
(Self::Buffered(filtered_output_writer), RecordProperties::Realigned { .. }) => {
filtered_output_writer.write(record, true)?;
}
}
Ok(())
}
}
pub struct FilteredOutputWriter {
buf: VecDeque<OutputRecord>,
out: bcf::Writer,
last_keep_pos: Option<i64>,
current_rid: Option<u32>,
current_pos: i64,
max_distance: i64,
}
impl FilteredOutputWriter {
fn write(
&mut self,
record: OutputRecord,
keep: bool,
) -> Result<(), rust_htslib::errors::Error> {
let rid = record.rid();
if rid != self.current_rid {
self.buf.clear();
self.last_keep_pos = None;
self.current_rid = rid;
}
self.current_pos = record.pos();
if keep {
self.last_keep_pos = Some(self.current_pos);
self.flush_relevant_from_buf()?;
}
match self.last_keep_pos {
Some(ts_pos) if self.current_pos < ts_pos + self.max_distance => {
self.out.write(&record)?;
}
Some(_) => {
self.last_keep_pos = None;
self.buf.push_back(record);
}
None => {
self.discard_stale_records();
self.buf.push_back(record);
}
}
Ok(())
}
fn flush_relevant_from_buf(&mut self) -> Result<(), rust_htslib::errors::Error> {
let Some(ts_pos) = self.last_keep_pos else {
return Ok(());
};
while let Some(front) = self.buf.front() {
if front.pos() < ts_pos - self.max_distance {
self.buf.pop_front(); } else {
let Some(record) = self.buf.pop_front() else {
break;
};
self.out.write(&record)?;
}
}
Ok(())
}
fn discard_stale_records(&mut self) {
let threshold = self.current_pos.saturating_sub(self.max_distance);
while let Some(front) = self.buf.front() {
if front.pos() < threshold {
self.buf.pop_front();
} else {
break;
}
}
}
}