use std::collections::HashSet;
use std::error::Error;
use noodles_sam as sam;
use rand::rngs::SmallRng;
use rand::{RngExt, SeedableRng};
use crate::cli::{OutputFormat, Sampler};
use crate::record::{DataRecord, from_alignment};
use crate::writer::OutputWriter;
pub struct Cfg {
pub output: Option<String>,
pub format: OutputFormat,
pub delimiter: String,
pub no_quotes: bool,
pub limit: usize,
pub detect_limit: usize,
pub keep_tag_prefix: bool,
pub sampler: Option<Sampler>,
pub seed: Option<u64>,
}
fn make_rng(seed: Option<u64>) -> SmallRng {
match seed {
Some(s) => SmallRng::seed_from_u64(s),
None => {
let mut sys = rand::rngs::SysRng;
SmallRng::try_from_rng(&mut sys).expect("system RNG unavailable")
}
}
}
pub fn run<Rec>(
records: &mut dyn Iterator<Item = std::io::Result<Rec>>,
header: &sam::Header,
cfg: &Cfg,
) -> Result<usize, Box<dyn Error>>
where
Rec: sam::alignment::Record,
{
match cfg.sampler {
Some(Sampler::Reservoir(n)) => run_reservoir(records, header, cfg, n),
Some(Sampler::Bernoulli(f)) => {
run_stream(records, header, cfg, Some((f, make_rng(cfg.seed))))
}
None => run_stream(records, header, cfg, None),
}
}
fn run_stream<Rec>(
records: &mut dyn Iterator<Item = std::io::Result<Rec>>,
header: &sam::Header,
cfg: &Cfg,
mut bernoulli: Option<(f64, SmallRng)>,
) -> Result<usize, Box<dyn Error>>
where
Rec: sam::alignment::Record,
{
let scan_cap = cfg.detect_limit.min(cfg.limit);
eprintln!("Scanning up to {scan_cap} records to discover all tags...");
let mut all_tags: HashSet<String> = HashSet::new();
let mut buffered: Vec<DataRecord> = Vec::with_capacity(scan_cap.min(1024));
for _ in 0..scan_cap {
match records.next() {
Some(Ok(record)) => {
let data = from_alignment(&record, header, cfg.keep_tag_prefix)?;
for tag in data.optional_fields.keys() {
all_tags.insert(tag.clone());
}
buffered.push(data);
}
Some(Err(err)) => return Err(err.into()),
None => break,
}
}
eprintln!(
"Found {} unique tags in {} scanned records.",
all_tags.len(),
buffered.len()
);
let mut sorted_tags: Vec<String> = all_tags.into_iter().collect();
sorted_tags.sort_unstable();
let mut writer = OutputWriter::new(
cfg.output.clone(),
cfg.format,
cfg.delimiter.clone(),
cfg.no_quotes,
&sorted_tags,
)?;
writer.write_header()?;
let mut written = 0usize;
let mut read = buffered.len();
for record in &buffered {
if let Some((fraction, rng)) = bernoulli.as_mut()
&& !rng.random_bool(*fraction)
{
continue;
}
writer.write_record(record)?;
written += 1;
}
eprintln!("Processing remaining records from stream...");
while read < cfg.limit {
match records.next() {
Some(Ok(record)) => {
read += 1;
let data = from_alignment(&record, header, cfg.keep_tag_prefix)?;
if let Some((fraction, rng)) = bernoulli.as_mut()
&& !rng.random_bool(*fraction)
{
continue;
}
writer.write_record(&data)?;
written += 1;
}
Some(Err(err)) => return Err(err.into()),
None => break,
}
}
writer.flush()?;
Ok(written)
}
fn run_reservoir<Rec>(
records: &mut dyn Iterator<Item = std::io::Result<Rec>>,
header: &sam::Header,
cfg: &Cfg,
n: usize,
) -> Result<usize, Box<dyn Error>>
where
Rec: sam::alignment::Record,
{
eprintln!("Reservoir-sampling {n} records (reading entire stream before writing)...");
let mut rng = make_rng(cfg.seed);
let mut all_tags: HashSet<String> = HashSet::new();
let mut reservoir: Vec<(usize, DataRecord)> = Vec::with_capacity(n);
let mut seen: usize = 0;
loop {
if cfg.limit != 0 && seen >= cfg.limit {
break; }
let record = match records.next() {
Some(Ok(record)) => record,
Some(Err(err)) => return Err(err.into()),
None => break,
};
seen += 1;
let data = from_alignment(&record, header, cfg.keep_tag_prefix)?;
for tag in data.optional_fields.keys() {
all_tags.insert(tag.clone());
}
if reservoir.len() < n {
reservoir.push((seen, data));
} else {
let j = rng.random_range(0..seen);
if j < n {
reservoir[j] = (seen, data);
}
}
}
let mut sorted_tags: Vec<String> = all_tags.into_iter().collect();
sorted_tags.sort_unstable();
let mut writer = OutputWriter::new(
cfg.output.clone(),
cfg.format,
cfg.delimiter.clone(),
cfg.no_quotes,
&sorted_tags,
)?;
writer.write_header()?;
reservoir.sort_unstable_by_key(|(position, _)| *position);
for (_, data) in &reservoir {
writer.write_record(data)?;
}
let written = reservoir.len();
writer.flush()?;
Ok(written)
}