use std::io;
use std::path::Path;
use std::time::{Duration, Instant};
use polars_core::prelude::DataFrame;
use polars_core::schema::Schema;
use polars_io::csv::write::CsvWriter;
use polars_io::prelude::SerWriter;
use polars_lazy::frame::{LazyCsvReader, LazyFileListReader, LazyFrame, ScanArgsParquet};
use super::error::{CommandError, Result};
use super::spec::{
BinaryInputs, IntervalColumns, IntervalCommandSpec, NearestCommandSpec, OverlapCommandSpec,
StrandConstraint, DEFAULT_CHROM_COL, DEFAULT_END_COL, DEFAULT_START_COL, DEFAULT_STRAND_COL,
};
use crate::{
DataFrameIntervalAccessors, GenomicNearestOptions, GenomicOverlapOptions, Strandedness,
};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ExecOptions {
pub delimiter: u8,
pub include_header: bool,
pub benchmark: bool,
pub reps: usize,
pub no_output: bool,
}
impl Default for ExecOptions {
fn default() -> Self {
Self {
delimiter: b'\t',
include_header: true,
benchmark: false,
reps: 1,
no_output: false,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum InputFormat {
Bed,
Csv,
Tsv,
Parquet,
}
pub fn run_command(spec: &IntervalCommandSpec, options: &ExecOptions) -> Result<()> {
if options.benchmark {
return run_benchmark_command(spec, options);
}
match spec {
IntervalCommandSpec::Overlap(spec) => run_overlap(spec, options),
IntervalCommandSpec::Nearest(spec) => run_nearest(spec, options),
}
}
pub fn execute_command_frame(spec: &IntervalCommandSpec) -> Result<DataFrame> {
match spec {
IntervalCommandSpec::Overlap(spec) => execute_overlap_frame(spec),
IntervalCommandSpec::Nearest(spec) => execute_nearest_frame(spec),
}
}
pub fn run_overlap(spec: &OverlapCommandSpec, options: &ExecOptions) -> Result<()> {
if options.benchmark {
return run_benchmark_command(&IntervalCommandSpec::Overlap(spec.clone()), options);
}
let mut frame = execute_overlap_frame(spec)?;
if options.no_output {
return Ok(());
}
write_frame(&mut frame, options)
}
pub fn run_nearest(spec: &NearestCommandSpec, options: &ExecOptions) -> Result<()> {
if options.benchmark {
return run_benchmark_command(&IntervalCommandSpec::Nearest(spec.clone()), options);
}
let mut frame = execute_nearest_frame(spec)?;
if options.no_output {
return Ok(());
}
write_frame(&mut frame, options)
}
pub fn execute_overlap_frame(spec: &OverlapCommandSpec) -> Result<DataFrame> {
let (left, right) =
load_and_validate_inputs(&spec.inputs, &spec.columns, spec.strand_constraint)?;
let options = overlap_options_from_spec(spec);
let left = left.collect()?;
let right = right.collect()?;
Ok(left.b().overlap(&right, options)?)
}
pub fn execute_nearest_frame(spec: &NearestCommandSpec) -> Result<DataFrame> {
let (left, right) =
load_and_validate_inputs(&spec.inputs, &spec.columns, spec.strand_constraint)?;
let options = nearest_options_from_spec(spec);
let left = left.collect()?;
let right = right.collect()?;
Ok(left.b().nearest(&right, options)?)
}
fn run_benchmark_command(spec: &IntervalCommandSpec, options: &ExecOptions) -> Result<()> {
let (operation, inputs, columns, strand_constraint) = command_parts(spec);
let started = Instant::now();
let (left_lazy, right_lazy) = load_and_validate_inputs(inputs, columns, strand_constraint)?;
let scan_validate = started.elapsed();
let started = Instant::now();
let left = left_lazy.collect()?;
let read_left = started.elapsed();
let started = Instant::now();
let right = right_lazy.collect()?;
let read_right = started.elapsed();
let reps = options.reps.max(1);
let mut operation_wall = Duration::ZERO;
let mut frame = None;
for _ in 0..reps {
let started = Instant::now();
let result = execute_preloaded_command(spec, &left, &right)?;
operation_wall += started.elapsed();
frame = Some(result);
}
let mut frame = frame.expect("benchmark repetitions are clamped to at least one");
eprintln!("result summary:");
eprintln!("{frame}");
let output_write = if options.no_output {
Duration::ZERO
} else {
let started = Instant::now();
write_frame(&mut frame, options)?;
started.elapsed()
};
let summary = BenchmarkSummary {
operation,
left_path: inputs.left_path.display().to_string(),
right_path: inputs.right_path.display().to_string(),
left_rows: left.height(),
right_rows: right.height(),
result_rows: frame.height(),
reps,
scan_validate,
read_left,
read_right,
operation_wall,
output_write,
no_output: options.no_output,
};
print_benchmark_summary(&summary);
Ok(())
}
fn execute_preloaded_command(
spec: &IntervalCommandSpec,
left: &DataFrame,
right: &DataFrame,
) -> Result<DataFrame> {
match spec {
IntervalCommandSpec::Overlap(spec) => {
Ok(left.b().overlap(right, overlap_options_from_spec(spec))?)
}
IntervalCommandSpec::Nearest(spec) => {
Ok(left.b().nearest(right, nearest_options_from_spec(spec))?)
}
}
}
fn overlap_options_from_spec(spec: &OverlapCommandSpec) -> GenomicOverlapOptions {
GenomicOverlapOptions {
strand_behavior: strand_behavior(spec.strand_constraint),
chromosome_col: spec.columns.chrom_col.clone(),
strand_col: spec.columns.resolved_strand_col(),
start_col: spec.columns.start_col.clone(),
end_col: spec.columns.end_col.clone(),
..GenomicOverlapOptions::default()
}
}
fn nearest_options_from_spec(spec: &NearestCommandSpec) -> GenomicNearestOptions {
GenomicNearestOptions {
suffix: spec.suffix.clone(),
distance_column: spec.distance_column.clone(),
preserve_input_order: spec.preserve_input_order,
strand_behavior: strand_behavior(spec.strand_constraint),
chromosome_col: spec.columns.chrom_col.clone(),
strand_col: spec.columns.resolved_strand_col(),
start_col: spec.columns.start_col.clone(),
end_col: spec.columns.end_col.clone(),
..GenomicNearestOptions::default()
}
}
fn command_parts(
spec: &IntervalCommandSpec,
) -> (
&'static str,
&BinaryInputs,
&IntervalColumns,
StrandConstraint,
) {
match spec {
IntervalCommandSpec::Overlap(spec) => (
"overlap",
&spec.inputs,
&spec.columns,
spec.strand_constraint,
),
IntervalCommandSpec::Nearest(spec) => (
"nearest",
&spec.inputs,
&spec.columns,
spec.strand_constraint,
),
}
}
struct BenchmarkSummary {
operation: &'static str,
left_path: String,
right_path: String,
left_rows: usize,
right_rows: usize,
result_rows: usize,
reps: usize,
scan_validate: Duration,
read_left: Duration,
read_right: Duration,
operation_wall: Duration,
output_write: Duration,
no_output: bool,
}
fn print_benchmark_summary(summary: &BenchmarkSummary) {
let operation_avg = summary.operation_wall / summary.reps as u32;
let measured_total = summary.scan_validate
+ summary.read_left
+ summary.read_right
+ summary.operation_wall
+ summary.output_write;
eprintln!();
eprintln!("polaranges benchmark");
eprintln!("====================");
eprintln!("{:<24} {}", "operation", summary.operation);
eprintln!("{:<24} {}", "left", summary.left_path);
eprintln!("{:<24} {}", "right", summary.right_path);
eprintln!("{:<24} {}", "left rows", summary.left_rows);
eprintln!("{:<24} {}", "right rows", summary.right_rows);
eprintln!("{:<24} {}", "result rows", summary.result_rows);
eprintln!("{:<24} {}", "repetitions", summary.reps);
eprintln!(
"{:<24} {}",
"stdout output",
output_status(summary.no_output)
);
eprintln!();
eprintln!("wall timings");
eprintln!("------------");
print_duration("scan + validate", summary.scan_validate);
print_duration("read left", summary.read_left);
print_duration("read right", summary.read_right);
print_duration("operation total", summary.operation_wall);
print_duration("operation avg", operation_avg);
print_duration("stdout write", summary.output_write);
print_duration("measured total", measured_total);
}
fn print_duration(label: &str, duration: Duration) {
eprintln!(" {:<22} {:>12.6} s", label, duration.as_secs_f64());
}
fn output_status(no_output: bool) -> &'static str {
if no_output {
"suppressed (--no-output)"
} else {
"written to stdout"
}
}
fn write_frame(frame: &mut DataFrame, options: &ExecOptions) -> Result<()> {
let stdout = io::stdout();
let handle = stdout.lock();
CsvWriter::new(handle)
.with_separator(options.delimiter)
.include_header(options.include_header)
.finish(frame)?;
Ok(())
}
fn load_and_validate_inputs(
inputs: &BinaryInputs,
columns: &IntervalColumns,
strand_constraint: StrandConstraint,
) -> Result<(LazyFrame, LazyFrame)> {
let left = scan_interval_file(&inputs.left_path)?;
let right = scan_interval_file(&inputs.right_path)?;
validate_required_columns(&inputs.left_path, &left, columns, strand_constraint)?;
validate_required_columns(&inputs.right_path, &right, columns, strand_constraint)?;
Ok((left, right))
}
fn strand_behavior(constraint: StrandConstraint) -> Strandedness {
match constraint {
StrandConstraint::Same => Strandedness::Same,
StrandConstraint::Opposite => Strandedness::Opposite,
StrandConstraint::Ignore => Strandedness::Ignore,
}
}
fn scan_interval_file(path: &Path) -> Result<LazyFrame> {
match detect_format(path)? {
InputFormat::Bed => scan_bed(path),
InputFormat::Csv => scan_delimited(path, b','),
InputFormat::Tsv => scan_delimited(path, b'\t'),
InputFormat::Parquet => scan_parquet(path),
}
}
fn scan_bed(path: &Path) -> Result<LazyFrame> {
let path_str = path.to_string_lossy();
let mut reader = LazyCsvReader::new(path_str.as_ref().into());
reader = reader.with_separator(b'\t').with_has_header(false);
Ok(reader.finish()?)
}
fn scan_delimited(path: &Path, separator: u8) -> Result<LazyFrame> {
let path_str = path.to_string_lossy();
let mut reader = LazyCsvReader::new(path_str.as_ref().into());
reader = reader.with_separator(separator);
Ok(reader.finish()?)
}
fn scan_parquet(path: &Path) -> Result<LazyFrame> {
let path_str = path.to_string_lossy();
Ok(LazyFrame::scan_parquet(
path_str.as_ref().into(),
ScanArgsParquet::default(),
)?)
}
fn detect_format(path: &Path) -> Result<InputFormat> {
let lower = path
.extension()
.and_then(|ext| ext.to_str())
.map(|ext| ext.to_ascii_lowercase())
.unwrap_or_default();
Ok(match lower.as_str() {
"bed" => InputFormat::Bed,
"csv" => InputFormat::Csv,
"tsv" => InputFormat::Tsv,
"parquet" | "pq" => InputFormat::Parquet,
_ => InputFormat::Tsv,
})
}
fn validate_required_columns(
path: &Path,
frame: &LazyFrame,
columns: &IntervalColumns,
strand_constraint: StrandConstraint,
) -> Result<()> {
let schema = collect_schema(frame)?;
let must_have = required_columns(columns, strand_constraint);
for column in must_have {
if !schema.contains(column) {
return Err(CommandError::MissingRequiredColumn {
path: path.to_path_buf(),
column: column.to_owned(),
});
}
}
Ok(())
}
fn collect_schema(frame: &LazyFrame) -> Result<Schema> {
let arc_schema = frame.clone().collect_schema()?;
Ok((*arc_schema).clone())
}
fn required_columns(
columns: &IntervalColumns,
strand_constraint: StrandConstraint,
) -> Vec<&'static str> {
let mut out = vec![DEFAULT_CHROM_COL, DEFAULT_START_COL, DEFAULT_END_COL];
if !matches!(strand_constraint, StrandConstraint::Ignore) {
out.push(DEFAULT_STRAND_COL);
}
let _ = columns;
out
}
#[cfg(test)]
mod tests {
}