use polars_core::prelude::{DataFrame, DataType, IdxCa};
use ruranges_core::{overlap_indices, overlaps};
use crate::error::{RangeFrameError, Result};
use crate::factorize::factorize_pair_by;
const DEFAULT_START_COL: &str = "Start";
const DEFAULT_END_COL: &str = "End";
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum OverlapMode {
All,
First,
Last,
}
impl OverlapMode {
fn as_kernel_str(self) -> &'static str {
match self {
Self::All => "all",
Self::First => "first",
Self::Last => "last",
}
}
}
#[derive(Clone, Debug)]
pub struct OverlapOptions {
pub multiple: OverlapMode,
pub slack: i64,
pub contained_intervals_only: bool,
pub preserve_input_order: bool,
pub match_by: Vec<String>,
pub left_start_col: String,
pub left_end_col: String,
pub right_start_col: String,
pub right_end_col: String,
}
impl Default for OverlapOptions {
fn default() -> Self {
Self {
multiple: OverlapMode::All,
slack: 0,
contained_intervals_only: false,
preserve_input_order: true,
match_by: Vec::new(),
left_start_col: DEFAULT_START_COL.to_owned(),
left_end_col: DEFAULT_END_COL.to_owned(),
right_start_col: DEFAULT_START_COL.to_owned(),
right_end_col: DEFAULT_END_COL.to_owned(),
}
}
}
impl OverlapOptions {
pub fn with_match_by<I, S>(mut self, columns: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
self.match_by = columns.into_iter().map(Into::into).collect();
self
}
pub fn with_interval_columns(
mut self,
start_col: impl Into<String>,
end_col: impl Into<String>,
) -> Self {
let start = start_col.into();
let end = end_col.into();
self.left_start_col = start.clone();
self.left_end_col = end.clone();
self.right_start_col = start;
self.right_end_col = end;
self
}
pub fn with_left_interval_columns(
mut self,
start_col: impl Into<String>,
end_col: impl Into<String>,
) -> Self {
self.left_start_col = start_col.into();
self.left_end_col = end_col.into();
self
}
pub fn with_right_interval_columns(
mut self,
start_col: impl Into<String>,
end_col: impl Into<String>,
) -> Self {
self.right_start_col = start_col.into();
self.right_end_col = end_col.into();
self
}
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
pub struct OverlapPairs {
pub left: Vec<u32>,
pub right: Vec<u32>,
}
impl OverlapPairs {
pub fn len(&self) -> usize {
self.left.len()
}
pub fn is_empty(&self) -> bool {
self.left.is_empty()
}
}
#[derive(Clone, Copy, Debug)]
pub(crate) struct RangeView<'a> {
pub(crate) df: &'a DataFrame,
pub(crate) start_col: &'a str,
pub(crate) end_col: &'a str,
}
impl<'a> RangeView<'a> {
pub(crate) fn new(df: &'a DataFrame) -> Result<Self> {
Self::with_columns(df, DEFAULT_START_COL, DEFAULT_END_COL)
}
pub(crate) fn with_columns(
df: &'a DataFrame,
start_col: &'a str,
end_col: &'a str,
) -> Result<Self> {
ensure_row_capacity(df.height())?;
ensure_coordinate_columns(df, start_col, end_col)?;
Ok(Self {
df,
start_col,
end_col,
})
}
pub(crate) fn overlap_pairs(
&self,
other: &RangeView<'_>,
options: &OverlapOptions,
) -> Result<OverlapPairs> {
self.overlap_pairs_by_match(other, &options.match_by, &options.match_by, options)
}
pub(crate) fn overlap_pairs_by_match(
&self,
other: &RangeView<'_>,
left_match_by: &[String],
right_match_by: &[String],
options: &OverlapOptions,
) -> Result<OverlapPairs> {
let (left_groups, right_groups) =
factorize_pair_by(self.df, other.df, left_match_by, right_match_by)?;
let left_coords = extract_coordinates(self.df, self.start_col, self.end_col)?;
let right_coords = extract_coordinates(other.df, other.start_col, other.end_col)?;
if let (Some((left_starts, left_ends)), Some((right_starts, right_ends))) =
(left_coords.as_i32(), right_coords.as_i32())
{
if let Ok(slack) = i32::try_from(options.slack) {
let (left, right) = overlaps::overlaps(
&left_groups,
left_starts,
left_ends,
&right_groups,
right_starts,
right_ends,
slack,
options.multiple.as_kernel_str(),
options.preserve_input_order,
options.contained_intervals_only,
);
return Ok(OverlapPairs { left, right });
}
}
if let (Some((left_starts, left_ends)), Some((right_starts, right_ends))) =
(left_coords.as_i64(), right_coords.as_i64())
{
let (left, right) = overlaps::overlaps(
&left_groups,
left_starts,
left_ends,
&right_groups,
right_starts,
right_ends,
options.slack,
options.multiple.as_kernel_str(),
options.preserve_input_order,
options.contained_intervals_only,
);
return Ok(OverlapPairs { left, right });
}
let (left_starts, left_ends) = left_coords.into_i64_buffers();
let (right_starts, right_ends) = right_coords.into_i64_buffers();
let (left, right) = overlaps::overlaps(
&left_groups,
&left_starts,
&left_ends,
&right_groups,
&right_starts,
&right_ends,
options.slack,
options.multiple.as_kernel_str(),
options.preserve_input_order,
options.contained_intervals_only,
);
Ok(OverlapPairs { left, right })
}
#[cfg(test)]
pub(crate) fn overlap(
&self,
other: &RangeView<'_>,
options: &OverlapOptions,
) -> Result<DataFrame> {
let indices =
self.overlap_indices_by_match(other, &options.match_by, &options.match_by, options)?;
take_rows_owned(self.df, indices)
}
pub(crate) fn overlap_indices_by_match(
&self,
other: &RangeView<'_>,
left_match_by: &[String],
right_match_by: &[String],
options: &OverlapOptions,
) -> Result<Vec<u32>> {
let (left_groups, right_groups) =
factorize_pair_by(self.df, other.df, left_match_by, right_match_by)?;
let left_coords = extract_coordinates(self.df, self.start_col, self.end_col)?;
let right_coords = extract_coordinates(other.df, other.start_col, other.end_col)?;
if let (Some((left_starts, left_ends)), Some((right_starts, right_ends))) =
(left_coords.as_i32(), right_coords.as_i32())
{
if let Ok(slack) = i32::try_from(options.slack) {
return Ok(overlap_indices::overlap_indices(
&left_groups,
left_starts,
left_ends,
&right_groups,
right_starts,
right_ends,
slack,
options.multiple.as_kernel_str(),
options.preserve_input_order,
options.contained_intervals_only,
));
}
}
if let (Some((left_starts, left_ends)), Some((right_starts, right_ends))) =
(left_coords.as_i64(), right_coords.as_i64())
{
return Ok(overlap_indices::overlap_indices(
&left_groups,
left_starts,
left_ends,
&right_groups,
right_starts,
right_ends,
options.slack,
options.multiple.as_kernel_str(),
options.preserve_input_order,
options.contained_intervals_only,
));
}
let (left_starts, left_ends) = left_coords.into_i64_buffers();
let (right_starts, right_ends) = right_coords.into_i64_buffers();
Ok(overlap_indices::overlap_indices(
&left_groups,
&left_starts,
&left_ends,
&right_groups,
&right_starts,
&right_ends,
options.slack,
options.multiple.as_kernel_str(),
options.preserve_input_order,
options.contained_intervals_only,
))
}
}
pub(crate) fn take_rows(df: &DataFrame, indices: &[u32]) -> Result<DataFrame> {
take_rows_owned(df, indices.to_vec())
}
pub(crate) fn take_rows_owned(df: &DataFrame, indices: Vec<u32>) -> Result<DataFrame> {
let idx = IdxCa::from_vec("idx".into(), indices);
Ok(df.take(&idx)?)
}
fn ensure_coordinate_columns(df: &DataFrame, start_col: &str, end_col: &str) -> Result<()> {
if df.column(start_col).is_err() {
return Err(RangeFrameError::MissingColumn(start_col.to_owned()));
}
if df.column(end_col).is_err() {
return Err(RangeFrameError::MissingColumn(end_col.to_owned()));
}
Ok(())
}
pub(crate) enum PreparedCoordinates<'a> {
I32 {
starts: NumericSlice<'a, i32>,
ends: NumericSlice<'a, i32>,
},
I64 {
starts: NumericSlice<'a, i64>,
ends: NumericSlice<'a, i64>,
},
}
impl<'a> PreparedCoordinates<'a> {
pub(crate) fn as_i32(&self) -> Option<(&[i32], &[i32])> {
match self {
Self::I32 { starts, ends } => Some((starts.as_slice(), ends.as_slice())),
Self::I64 { .. } => None,
}
}
pub(crate) fn as_i64(&self) -> Option<(&[i64], &[i64])> {
match self {
Self::I32 { .. } => None,
Self::I64 { starts, ends } => Some((starts.as_slice(), ends.as_slice())),
}
}
pub(crate) fn into_i64_buffers(self) -> (Vec<i64>, Vec<i64>) {
match self {
Self::I32 { starts, ends } => (
starts
.as_slice()
.iter()
.map(|value| *value as i64)
.collect(),
ends.as_slice().iter().map(|value| *value as i64).collect(),
),
Self::I64 { starts, ends } => (starts.into_vec(), ends.into_vec()),
}
}
}
pub(crate) fn extract_coordinates<'a>(
df: &'a DataFrame,
start_col: &str,
end_col: &str,
) -> Result<PreparedCoordinates<'a>> {
let start_dtype = df.column(start_col)?.dtype().clone();
let end_dtype = df.column(end_col)?.dtype().clone();
if matches!(
(&start_dtype, &end_dtype),
(DataType::Int32, DataType::Int32)
) {
let starts = extract_i32_slice(df, start_col)?;
let ends = extract_i32_slice(df, end_col)?;
validate_i32_coordinates(starts.as_slice(), ends.as_slice(), start_col, end_col)?;
return Ok(PreparedCoordinates::I32 { starts, ends });
}
if matches!(
(&start_dtype, &end_dtype),
(DataType::Int64, DataType::Int64)
) {
let starts = extract_i64_slice(df, start_col)?;
let ends = extract_i64_slice(df, end_col)?;
validate_i64_coordinates(starts.as_slice(), ends.as_slice(), start_col, end_col)?;
return Ok(PreparedCoordinates::I64 { starts, ends });
}
let starts = extract_i64_column(df, start_col)?;
let ends = extract_i64_column(df, end_col)?;
validate_i64_coordinates(&starts, &ends, start_col, end_col)?;
Ok(PreparedCoordinates::I64 {
starts: NumericSlice::Owned(starts),
ends: NumericSlice::Owned(ends),
})
}
fn validate_i32_coordinates(
starts: &[i32],
ends: &[i32],
start_col: &str,
end_col: &str,
) -> Result<()> {
if starts.iter().any(|value| *value < 0) {
return Err(RangeFrameError::NegativeCoordinate {
column: start_col.to_owned(),
});
}
if ends.iter().any(|value| *value < 0) {
return Err(RangeFrameError::NegativeCoordinate {
column: end_col.to_owned(),
});
}
if starts.iter().zip(ends).any(|(start, end)| start >= end) {
return Err(RangeFrameError::InvalidIntervals {
start: start_col.to_owned(),
end: end_col.to_owned(),
});
}
Ok(())
}
fn validate_i64_coordinates(
starts: &[i64],
ends: &[i64],
start_col: &str,
end_col: &str,
) -> Result<()> {
if starts.iter().any(|value| *value < 0) {
return Err(RangeFrameError::NegativeCoordinate {
column: start_col.to_owned(),
});
}
if ends.iter().any(|value| *value < 0) {
return Err(RangeFrameError::NegativeCoordinate {
column: end_col.to_owned(),
});
}
if starts.iter().zip(ends).any(|(start, end)| start >= end) {
return Err(RangeFrameError::InvalidIntervals {
start: start_col.to_owned(),
end: end_col.to_owned(),
});
}
Ok(())
}
fn extract_i64_column(df: &DataFrame, column_name: &str) -> Result<Vec<i64>> {
let column = df.column(column_name)?;
let series = column.as_materialized_series();
if column.null_count() > 0 {
return Err(RangeFrameError::NullValues {
column: column_name.to_owned(),
});
}
let casted = match series.dtype() {
DataType::Int8
| DataType::Int16
| DataType::Int32
| DataType::Int64
| DataType::UInt8
| DataType::UInt16
| DataType::UInt32
| DataType::UInt64 => series.cast(&DataType::Int64)?,
dtype => {
return Err(RangeFrameError::InvalidCoordinateDtype {
column: column_name.to_owned(),
dtype: dtype.to_string(),
});
}
};
Ok(casted.i64()?.into_no_null_iter().collect())
}
pub(crate) enum NumericSlice<'a, T> {
Borrowed(&'a [T]),
Owned(Vec<T>),
}
impl<'a, T> NumericSlice<'a, T> {
fn as_slice(&self) -> &[T] {
match self {
Self::Borrowed(slice) => slice,
Self::Owned(values) => values.as_slice(),
}
}
fn into_vec(self) -> Vec<T>
where
T: Clone,
{
match self {
Self::Borrowed(slice) => slice.to_vec(),
Self::Owned(values) => values,
}
}
}
fn extract_i32_slice<'a>(df: &'a DataFrame, column_name: &str) -> Result<NumericSlice<'a, i32>> {
let column = df.column(column_name)?;
if column.null_count() > 0 {
return Err(RangeFrameError::NullValues {
column: column_name.to_owned(),
});
}
let series = column.as_materialized_series();
let ca = series
.i32()
.map_err(|_| RangeFrameError::InvalidCoordinateDtype {
column: column_name.to_owned(),
dtype: series.dtype().to_string(),
})?;
if let Ok(slice) = ca.cont_slice() {
return Ok(NumericSlice::Borrowed(slice));
}
Ok(NumericSlice::Owned(ca.into_no_null_iter().collect()))
}
fn extract_i64_slice<'a>(df: &'a DataFrame, column_name: &str) -> Result<NumericSlice<'a, i64>> {
let column = df.column(column_name)?;
if column.null_count() > 0 {
return Err(RangeFrameError::NullValues {
column: column_name.to_owned(),
});
}
let series = column.as_materialized_series();
let ca = series
.i64()
.map_err(|_| RangeFrameError::InvalidCoordinateDtype {
column: column_name.to_owned(),
dtype: series.dtype().to_string(),
})?;
if let Ok(slice) = ca.cont_slice() {
return Ok(NumericSlice::Borrowed(slice));
}
Ok(NumericSlice::Owned(ca.into_no_null_iter().collect()))
}
fn ensure_row_capacity(len: usize) -> Result<()> {
if len > u32::MAX as usize {
return Err(RangeFrameError::TooManyRows { len });
}
Ok(())
}
#[cfg(test)]
mod tests {
use polars_core::prelude::{Column, DataFrame, NamedFrom, Series};
use super::{take_rows, OverlapMode, OverlapOptions, RangeView};
use crate::RangeFrameError;
fn make_df(columns: Vec<Series>) -> DataFrame {
let columns = columns.into_iter().map(Column::from).collect::<Vec<_>>();
DataFrame::new_infer_height(columns).unwrap()
}
#[test]
fn overlap_pairs_follow_kernel_row_ids() {
let left_df = make_df(vec![
Series::new("Start".into(), &[1_i64, 10, 30]),
Series::new("End".into(), &[5_i64, 20, 40]),
Series::new("Chrom".into(), &["chr1", "chr1", "chr1"]),
]);
let right_df = make_df(vec![
Series::new("Start".into(), &[3_i64, 11, 18, 35]),
Series::new("End".into(), &[4_i64, 12, 19, 36]),
Series::new("Chrom".into(), &["chr1", "chr1", "chr1", "chr1"]),
]);
let left = RangeView::new(&left_df).unwrap();
let right = RangeView::new(&right_df).unwrap();
let options = OverlapOptions::default().with_match_by(["Chrom"]);
let pairs = left.overlap_pairs(&right, &options).unwrap();
assert_eq!(pairs.left, vec![0, 1, 1, 2]);
assert_eq!(pairs.right, vec![0, 1, 2, 3]);
}
#[test]
fn overlap_gathers_left_rows() {
let left_df = make_df(vec![
Series::new("Start".into(), &[1_i64, 10, 30]),
Series::new("End".into(), &[5_i64, 20, 40]),
]);
let right_df = make_df(vec![
Series::new("Start".into(), &[3_i64, 11, 18, 35]),
Series::new("End".into(), &[4_i64, 12, 19, 36]),
]);
let left = RangeView::new(&left_df).unwrap();
let right = RangeView::new(&right_df).unwrap();
let result = left.overlap(&right, &OverlapOptions::default()).unwrap();
assert_eq!(result.height(), 4);
assert_eq!(
result
.column("Start")
.unwrap()
.i64()
.unwrap()
.into_no_null_iter()
.collect::<Vec<_>>(),
vec![1, 10, 10, 30]
);
}
#[test]
fn overlap_respects_multiple_first() {
let left_df = make_df(vec![
Series::new("Start".into(), &[10_i64]),
Series::new("End".into(), &[20_i64]),
]);
let right_df = make_df(vec![
Series::new("Start".into(), &[11_i64, 12]),
Series::new("End".into(), &[13_i64, 14]),
]);
let left = RangeView::new(&left_df).unwrap();
let right = RangeView::new(&right_df).unwrap();
let options = OverlapOptions {
multiple: OverlapMode::First,
..OverlapOptions::default()
};
let pairs = left.overlap_pairs(&right, &options).unwrap();
assert_eq!(pairs.left, vec![0]);
assert_eq!(pairs.right, vec![0]);
}
#[test]
fn overlap_pairs_support_i32_columns_without_casting_to_i64_first() {
let left_df = make_df(vec![
Series::new("Start".into(), &[1_i32, 10, 30]),
Series::new("End".into(), &[5_i32, 20, 40]),
Series::new("Chrom".into(), &["chr1", "chr1", "chr1"]),
]);
let right_df = make_df(vec![
Series::new("Start".into(), &[3_i32, 11, 18, 35]),
Series::new("End".into(), &[4_i32, 12, 19, 36]),
Series::new("Chrom".into(), &["chr1", "chr1", "chr1", "chr1"]),
]);
let left = RangeView::new(&left_df).unwrap();
let right = RangeView::new(&right_df).unwrap();
let options = OverlapOptions::default().with_match_by(["Chrom"]);
let pairs = left.overlap_pairs(&right, &options).unwrap();
assert_eq!(pairs.left, vec![0, 1, 1, 2]);
assert_eq!(pairs.right, vec![0, 1, 2, 3]);
}
#[test]
fn take_rows_uses_row_ids_from_overlap_pairs() {
let df = make_df(vec![
Series::new("Start".into(), &[1_i64, 10, 30]),
Series::new("End".into(), &[5_i64, 20, 40]),
]);
let taken = take_rows(&df, &[2, 0]).unwrap();
assert_eq!(
taken
.column("Start")
.unwrap()
.i64()
.unwrap()
.into_no_null_iter()
.collect::<Vec<_>>(),
vec![30, 1]
);
}
#[test]
fn overlap_rejects_invalid_intervals_without_pre_materializing_twice() {
let left_df = make_df(vec![
Series::new("Start".into(), &[10_i32]),
Series::new("End".into(), &[10_i32]),
]);
let right_df = make_df(vec![
Series::new("Start".into(), &[1_i32]),
Series::new("End".into(), &[2_i32]),
]);
let left = RangeView::new(&left_df).unwrap();
let right = RangeView::new(&right_df).unwrap();
let error = left
.overlap_pairs(&right, &OverlapOptions::default())
.unwrap_err();
assert!(matches!(error, RangeFrameError::InvalidIntervals { .. }));
}
}