use std::io::{self, BufRead, BufReader, Write};
use std::process;
use clap::Parser;
use poa_consensus::{AlignmentMode, DiagnoseConfig, PoaConfig, PoaError, auto_orient, diagnose};
#[derive(Parser)]
#[command(name = "poa-consensus", version)]
struct Args {
#[arg(default_value = "-")]
input: String,
#[arg(short, long)]
seed: Option<usize>,
#[arg(short = 'b', long, default_value_t = 0)]
band_width: usize,
#[arg(long)]
adaptive_band: bool,
#[arg(long, default_value_t = 3)]
min_reads: usize,
#[arg(short = 'm', long)]
multi: bool,
#[arg(long)]
semi_global: bool,
#[arg(short = 'q', long)]
quiet: bool,
}
fn main() {
if let Err(e) = run() {
eprintln!("error: {e}");
process::exit(1);
}
}
fn run() -> Result<(), Box<dyn std::error::Error>> {
let args = Args::parse();
let reads = load_reads(&args.input)?;
if reads.is_empty() {
return Err("no reads in input".into());
}
for seq in &reads {
if seq
.iter()
.any(|&b| !matches!(b, b'A' | b'C' | b'G' | b'T' | b'a' | b'c' | b'g' | b't'))
{
eprintln!(
"poa-consensus: warning: read contains non-ACGT bases; \
treating as mismatches"
);
break;
}
}
if reads.len() == 1 {
let stdout = io::stdout();
let mut out = stdout.lock();
writeln!(out, ">consensus reads=1 seed=0 band=unbanded")?;
out.write_all(&reads[0])?;
writeln!(out)?;
return Ok(());
}
let seed_idx = match args.seed {
Some(idx) => {
if idx >= reads.len() {
return Err(format!("seed {} out of range ({} reads)", idx, reads.len()).into());
}
idx
}
None => median_seed(&reads),
};
let oriented = auto_orient(&reads, seed_idx);
let slices: Vec<&[u8]> = oriented.iter().map(|r| r.as_ref()).collect();
let config = PoaConfig {
band_width: args.band_width,
adaptive_band: args.adaptive_band,
min_reads: args.min_reads,
alignment_mode: if args.semi_global {
AlignmentMode::SemiGlobal
} else {
AlignmentMode::Global
},
..PoaConfig::default()
};
let band_desc = if args.adaptive_band {
"adaptive".to_string()
} else if args.band_width == 0 {
"unbanded".to_string()
} else {
args.band_width.to_string()
};
let stdout = io::stdout();
let mut out = stdout.lock();
let n = reads.len();
if args.multi {
let alleles = poa_consensus::consensus_multi(&slices, seed_idx, &config)
.inspect_err(|e| explain_error(e, n))?;
let total = alleles.len();
let allele_cfg = DiagnoseConfig {
is_allele_partition: true,
..DiagnoseConfig::default()
};
for (i, allele) in alleles.iter().enumerate() {
if !args.quiet {
let label = format!("allele {}/{}", i + 1, total);
emit_warnings(&diagnose(allele, &allele_cfg), &label);
}
writeln!(
out,
">allele_{} reads={} total_reads={n} seed={seed_idx} band={band_desc} allele={}/{}",
i + 1,
allele.n_reads,
i + 1,
total
)?;
out.write_all(&allele.sequence)?;
writeln!(out)?;
}
} else {
let was_banded = config.band_width > 0 || config.adaptive_band;
let mut result = poa_consensus::consensus(&slices, seed_idx, &config)
.inspect_err(|e| explain_error(e, n))?;
if was_banded {
let diag = diagnose(&result, &DiagnoseConfig::default());
if let Some(ref t) = diag.truncation_suspected {
if t.median_read_len <= 5_000 {
let mut cfg2 = config.clone();
cfg2.band_width = 0;
cfg2.adaptive_band = false;
if let Ok(r2) = poa_consensus::consensus(&slices, seed_idx, &cfg2) {
result = r2;
}
} else if !args.quiet {
eprintln!(
"poa-consensus: warning: consensus ({} bp) is {:.0}% of median \
read length ({} bp) — suspected banded DP truncation on a long \
read set; retry with --band-width 0 to correct",
t.consensus_len,
t.ratio * 100.0,
t.median_read_len,
);
}
}
}
if !args.quiet {
emit_warnings(&diagnose(&result, &DiagnoseConfig::default()), "consensus");
}
writeln!(out, ">consensus reads={n} seed={seed_idx} band={band_desc}")?;
out.write_all(&result.sequence)?;
writeln!(out)?;
}
Ok(())
}
fn explain_error(e: &PoaError, n_reads: usize) {
match e {
PoaError::InsufficientDepth { got, min } => {
eprintln!("poa-consensus: error: only {got} read(s) provided, minimum is {min}");
if *got > 0 {
eprintln!(
" hint: use --min-reads {got} to lower the floor (accuracy will \
suffer at low depth)"
);
}
}
PoaError::BandTooNarrow {
configured,
required,
} => {
eprintln!(
"poa-consensus: error: band width {configured} is too narrow \
(need ≥ {required} for this read set)"
);
eprintln!(" hint: try --adaptive-band, or --band-width {required}");
}
PoaError::NoSpanningReads {
left_depth,
right_depth,
} => {
eprintln!(
"poa-consensus: error: no read spans the full locus \
({left_depth} left-only, {right_depth} right-only reads)"
);
eprintln!(
" hint: if reads are split into two non-overlapping groups, \
use bridged_consensus() to assemble each side separately"
);
}
_ => {} }
let _ = n_reads; }
fn emit_warnings(warnings: &poa_consensus::ConsensusWarnings, label: &str) {
for (is_warning, msg) in warnings.messages(label) {
let level = if is_warning { "warning" } else { "note" };
eprintln!("poa-consensus: {level}: {msg}");
}
}
fn load_reads(path: &str) -> Result<Vec<Vec<u8>>, Box<dyn std::error::Error>> {
if path == "-" {
let stdin = io::stdin();
let mut buf = BufReader::new(stdin.lock());
parse_reads(&mut buf)
} else {
let file = std::fs::File::open(path)?;
let mut buf = BufReader::new(file);
parse_reads(&mut buf)
}
}
fn parse_reads<R: BufRead>(reader: &mut R) -> Result<Vec<Vec<u8>>, Box<dyn std::error::Error>> {
let first = first_byte(reader)?;
match first {
Some(b'>') => parse_fasta(reader),
Some(b'@') => parse_fastq(reader),
Some(b) => Err(format!(
"unexpected first byte 0x{b:02x}; expected '>' (FASTA) or '@' (FASTQ)"
)
.into()),
None => Ok(vec![]),
}
}
fn parse_fasta<R: BufRead>(reader: &mut R) -> Result<Vec<Vec<u8>>, Box<dyn std::error::Error>> {
use noodles::fasta;
let mut fa_reader = fasta::io::Reader::new(reader);
let mut reads = Vec::new();
for result in fa_reader.records() {
let record = result?;
reads.push(record.sequence().as_ref().to_vec());
}
Ok(reads)
}
fn parse_fastq<R: BufRead>(reader: &mut R) -> Result<Vec<Vec<u8>>, Box<dyn std::error::Error>> {
use noodles::fastq;
let mut fq_reader = fastq::io::Reader::new(reader);
let mut reads = Vec::new();
for result in fq_reader.records() {
let record = result?;
reads.push(record.sequence().to_vec());
}
Ok(reads)
}
fn first_byte<R: BufRead>(reader: &mut R) -> Result<Option<u8>, Box<dyn std::error::Error>> {
loop {
let buf = reader.fill_buf()?;
if buf.is_empty() {
return Ok(None);
}
let b = buf[0];
if b == b'\n' || b == b'\r' || b == b' ' {
reader.consume(1);
continue;
}
return Ok(Some(b));
}
}
fn median_seed(reads: &[Vec<u8>]) -> usize {
let mut order: Vec<usize> = (0..reads.len()).collect();
order.sort_unstable_by_key(|&i| reads[i].len());
order[order.len() / 2]
}