mod output;
mod stats;
mod utils;
use std::path::Path;
use std::sync::Arc;
use anyhow::Result;
use log::{debug, info, warn};
use ncbi_vdb_sys::SraReader;
use output::{build_segment_writer, BoxedSegmentWriter};
use parking_lot::Mutex;
use crate::cli::{DumpOutput, FilterOptions, InputOptions, OutputFormat};
use crate::output::{build_path_name, OutputFileType};
use crate::prefetch::identify_url;
use crate::RECORD_CAPACITY;
use crate::utils::get_num_records;
use stats::ProcessStatistics;
use utils::write_segment_to_buffer_set;
fn launch_threads(
path: &str,
num_threads: u64,
records_per_thread: u64,
remainder: u64,
writer: Arc<Mutex<BoxedSegmentWriter>>,
filter_opts: FilterOptions,
format: OutputFormat,
accession_prefix: Option<String>,
include_sid: bool,
) -> Result<ProcessStatistics> {
let segment_set = if filter_opts.include.is_empty() {
None
} else {
let set: Vec<usize> = filter_opts.include.clone();
Some(set)
};
let mut handles = Vec::new();
for i in 0..num_threads {
let segment_set = segment_set.clone();
let start = (i * records_per_thread) + 1;
let stop = if i == num_threads - 1 {
start + records_per_thread + remainder - 1
} else {
start + records_per_thread - 1
};
let path = path.to_string();
let shared_writer = writer.clone();
let accession_prefix = accession_prefix.clone();
let handle = std::thread::spawn(move || -> Result<ProcessStatistics> {
let reader = SraReader::new(&path)?;
let mut stats = ProcessStatistics::default();
let mut local_buffers = shared_writer.lock().generate_local_buffers();
let mut counts = vec![0; local_buffers.len()];
for (idx, record) in reader.into_range_iter(start as i64, stop)?.enumerate() {
let record = record?;
for segment in record.into_iter() {
if let Some(ref set) = segment_set {
if !set.contains(&segment.sid()) {
continue;
}
}
if filter_opts.skip_technical && segment.is_technical() {
stats.inc_filter_type(segment.sid());
continue;
}
if segment.len() < filter_opts.min_read_len {
stats.inc_filter_size(segment.sid());
continue;
}
write_segment_to_buffer_set(
&mut local_buffers,
&segment,
format,
accession_prefix.as_deref(),
include_sid,
)?;
if counts.len() == 1 {
counts[0] += 1;
} else {
counts[segment.sid()] += 1;
}
stats.inc_reads(segment.sid());
}
if idx > 0 && (idx % RECORD_CAPACITY == 0) {
shared_writer
.lock()
.write_all_buffers(&mut local_buffers, &mut counts)?;
}
stats.inc_spots();
}
shared_writer
.lock()
.write_all_buffers(&mut local_buffers, &mut counts)?;
Ok(stats)
});
handles.push(handle);
}
let mut stats = ProcessStatistics::default();
for handle in handles {
let thread_stats = handle.join().expect("Thread panicked")?;
stats = stats + thread_stats;
}
Ok(stats)
}
pub fn dump(
input: &InputOptions,
num_threads: u64,
output_opts: &DumpOutput,
filter_opts: FilterOptions,
) -> Result<()> {
let accession = if !Path::new(&input.accession).exists() {
info!(accession = input.accession.as_str(); "Identifying SRA data URL for accession");
let runtime = tokio::runtime::Runtime::new()?;
let url = runtime.block_on(identify_url(&input.accession, &input.options))?;
info!(url = url.as_str(); "Streaming SRA records from URL");
url
} else {
debug!(path = input.accession.as_str(); "Using local SRA file");
input.accession.to_string()
};
let num_records = get_num_records(&accession)?;
let num_records = if let Some(limit) = filter_opts.limit {
if limit > num_records {
warn!(
spot_limit = limit,
actual_spots = num_records;
"Provided spot limit exceeds actual number of spots, processing full archive"
);
}
num_records.min(limit)
} else {
num_records
};
let records_per_thread = num_records / num_threads;
let remainder = num_records % num_threads;
let writer = build_segment_writer(
Some(&output_opts.outdir),
&output_opts.prefix,
output_opts.compression,
output_opts.format,
num_threads as usize,
&filter_opts,
output_opts.named_pipes,
output_opts.split,
)
.map(|x| Arc::new(Mutex::new(x)))?;
let included_segs = filter_opts.include.clone();
let accession_name = Path::new(&input.accession)
.file_stem()
.and_then(|s| s.to_str())
.unwrap_or(&input.accession)
.to_string();
let accession_prefix = Some(accession_name);
let stats = launch_threads(
&accession,
num_threads,
records_per_thread,
remainder,
writer,
filter_opts,
output_opts.format,
accession_prefix,
output_opts.include_sid(),
)?;
if output_opts.split {
let wrap = |x| {
if output_opts.named_pipes {
OutputFileType::NamedPipe(x)
} else {
OutputFileType::RegularFile(x)
}
};
stats.reads_per_segment.iter().enumerate().try_for_each(
|(seg_id, &count)| -> Result<()> {
if !included_segs.is_empty() && !included_segs.contains(&seg_id) {
return Ok(());
}
if count == 0 || output_opts.named_pipes {
let path = build_path_name(
wrap(&output_opts.outdir),
&output_opts.prefix,
output_opts.compression,
output_opts.format,
seg_id,
);
if output_opts.keep_empty {
warn!(path = path.as_str(), segment_id = seg_id; "Output file is empty but kept due to --keep-empty flag");
} else {
debug!(path = path.as_str(), segment_id = seg_id; "Removing empty output file");
std::fs::remove_file(path)?;
}
}
Ok(())
},
)?;
}
stats.pprint(&mut std::io::stderr())?;
Ok(())
}