#[cfg(feature = "bedclassifier")]
use crate::errors::BedClassifierError;
#[cfg(feature = "bedclassifier")]
use gtars_core::models::RegionSet;
#[cfg(feature = "bedclassifier")]
use regex::Regex;
use std::fmt::{self, Display};
#[cfg(feature = "bedclassifier")]
use polars::prelude::*;
#[derive(Clone, Debug, PartialEq)]
pub enum DataFormat {
Unknown,
UcscBed,
UcscBedRs,
BedLike,
BedLikeRs,
EncodeNarrowPeak,
EncodeNarrowPeakRs,
EncodeBroadPeak,
EncodeBroadPeakRs,
EncodeGappedPeak,
EncodeGappedPeakRs,
EncodeRnaElements,
EncodeRnaElementsRs,
}
impl Display for DataFormat {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let s = match self {
DataFormat::Unknown => "unknown_data_format",
DataFormat::UcscBed => "ucsc_bed",
DataFormat::UcscBedRs => "ucsc_bed_rs",
DataFormat::BedLike => "bed_like",
DataFormat::BedLikeRs => "bed_like_rs",
DataFormat::EncodeNarrowPeak => "encode_narrowpeak",
DataFormat::EncodeNarrowPeakRs => "encode_narrowpeak_rs",
DataFormat::EncodeBroadPeak => "encode_broadpeak",
DataFormat::EncodeBroadPeakRs => "encode_broadpeak_rs",
DataFormat::EncodeGappedPeak => "encode_gappedpeak",
DataFormat::EncodeGappedPeakRs => "encode_gappedpeak_rs",
DataFormat::EncodeRnaElements => "encode_rna_elements",
DataFormat::EncodeRnaElementsRs => "encode_rna_elements_rs",
};
write!(f, "{}", s)
}
}
#[derive(Clone, Debug)]
pub struct BedClassificationOutput {
pub bed_compliance: String,
pub data_format: DataFormat,
pub compliant_columns: usize,
pub non_compliant_columns: usize,
}
impl Display for BedClassificationOutput {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"BedClassificationOutput {{ bed_compliance: {}, data_format: {}, compliant_columns: {}, non_compliant_columns: {} }}",
self.bed_compliance, self.data_format, self.compliant_columns, self.non_compliant_columns
)
}
}
#[cfg(feature = "bedclassifier")]
pub fn classify_bed(region_set: &RegionSet) -> Result<BedClassificationOutput, BedClassifierError> {
let df = match region_set.to_polars() {
Ok(df) => df,
Err(_) => {
return Ok(BedClassificationOutput {
bed_compliance: "unknown_bed_compliance".to_string(),
data_format: DataFormat::Unknown,
compliant_columns: 0,
non_compliant_columns: 0,
});
}
};
let num_cols = df.width();
let mut compliant_columns = 0;
let mut relaxed = false;
let check_string_column = |df: &DataFrame, col_idx: usize, pattern: &str| -> bool {
if let Ok(col) = df.column(&format!("column_{}", col_idx + 1)) {
if let Ok(series) = col.cast(&DataType::String) {
let regex = Regex::new(pattern).unwrap();
if let Ok(ca) = series.str() {
return ca
.into_iter()
.all(|opt| opt.map(|s| regex.is_match(s)).unwrap_or(false));
}
}
}
false
};
let check_int_column =
|df: &DataFrame, col_idx: usize, min_val: Option<i64>, max_val: Option<i64>| -> bool {
if let Ok(col) = df.column(&format!("column_{}", col_idx + 1)) {
if let Ok(ca) = col.i64() {
return ca.into_iter().all(|opt| {
opt.map(|val| {
let mut valid = true;
if let Some(min) = min_val {
valid = valid && val >= min;
}
if let Some(max) = max_val {
valid = valid && val <= max;
}
valid
})
.unwrap_or(false)
});
}
}
false
};
let check_float_or_minus_one = |df: &DataFrame, col_idx: usize| -> bool {
if let Ok(col) = df.column(&format!("column_{}", col_idx + 1)) {
if col.dtype().is_float() {
return true;
}
if let Ok(ca) = col.i64() {
return ca
.into_iter()
.all(|opt| opt.map(|v| v == -1).unwrap_or(false));
}
}
false
};
let regex_colors =
r"^(?:\d|[1-9]\d|1\d{2}|2[0-4]\d|25[0-5])(?:,(?:\d|[1-9]\d|1\d{2}|2[0-4]\d|25[0-5])){0,2}$";
for col_idx in 0..num_cols {
let is_valid = match col_idx {
0 => check_string_column(&df, col_idx, r"[A-Za-z0-9_]{1,255}"),
1 => check_int_column(&df, col_idx, Some(0), None),
2 => check_int_column(&df, col_idx, Some(0), None),
3 => check_string_column(&df, col_idx, r"[\x20-\x7e]{1,255}"),
4 => {
if check_int_column(&df, col_idx, Some(0), Some(1000)) {
true
} else {
if check_int_column(&df, col_idx, Some(0), None) {
relaxed = true;
true
} else {
false
}
}
}
5 => {
if let Ok(col) = df.column(&format!("column_{}", col_idx + 1)) {
if let Ok(series) = col.cast(&DataType::String) {
if let Ok(ca) = series.str() {
ca.into_iter().all(|opt| {
opt.map(|s| s == "+" || s == "-" || s == ".")
.unwrap_or(false)
})
} else {
false
}
} else {
false
}
} else {
false
}
}
6 => check_int_column(&df, col_idx, Some(0), None),
7 => check_int_column(&df, col_idx, Some(0), None),
8 => check_string_column(&df, col_idx, regex_colors),
9 => check_int_column(&df, col_idx, None, None),
10 => check_string_column(&df, col_idx, r"^(0(,\d+)*|\d+(,\d+)*)?,?$"),
11 => check_string_column(&df, col_idx, r"^(0(,\d+)*|\d+(,\d+)*)?,?$"),
12 => check_float_or_minus_one(&df, col_idx),
13 => {
if let Ok(col) = df.column(&format!("column_{}", col_idx + 1)) {
if let Ok(ca) = col.i64() {
if let Some(first) = ca.get(0) {
first != -1
} else {
false
}
} else {
false
}
} else {
false
}
}
_ => false,
};
if is_valid && col_idx < 12 {
compliant_columns += 1;
} else {
let nccols = num_cols - compliant_columns;
if col_idx >= 6 {
if num_cols == 10
&& col_idx == 6
&& check_float_or_minus_one(&df, 6)
&& check_float_or_minus_one(&df, 7)
&& check_float_or_minus_one(&df, 8)
&& check_int_column(&df, 9, None, None)
{
return Ok(BedClassificationOutput {
bed_compliance: format!("bed{}+{}", compliant_columns, nccols),
data_format: if relaxed {
DataFormat::EncodeNarrowPeakRs
} else {
DataFormat::EncodeNarrowPeak
},
compliant_columns,
non_compliant_columns: nccols,
});
}
if num_cols == 9 && col_idx == 6 {
if check_float_or_minus_one(&df, 6)
&& check_float_or_minus_one(&df, 7)
&& check_float_or_minus_one(&df, 8)
{
return Ok(BedClassificationOutput {
bed_compliance: format!("bed{}+{}", compliant_columns, nccols),
data_format: if relaxed {
DataFormat::EncodeBroadPeakRs
} else {
DataFormat::EncodeBroadPeak
},
compliant_columns,
non_compliant_columns: nccols,
});
} else if check_float_or_minus_one(&df, 6) && check_float_or_minus_one(&df, 7) {
if let Ok(col) = df.column(&format!("column_{}", 9)) {
if let Ok(ca) = col.i64() {
if let Some(first) = ca.get(0) {
if first != -1 {
return Ok(BedClassificationOutput {
bed_compliance: format!(
"bed{}+{}",
compliant_columns, nccols
),
data_format: if relaxed {
DataFormat::EncodeRnaElementsRs
} else {
DataFormat::EncodeRnaElements
},
compliant_columns,
non_compliant_columns: nccols,
});
}
}
}
}
}
}
if num_cols == 15
&& col_idx == 12
&& check_float_or_minus_one(&df, 12)
&& check_float_or_minus_one(&df, 13)
&& check_float_or_minus_one(&df, 14)
{
return Ok(BedClassificationOutput {
bed_compliance: format!("bed{}+{}", compliant_columns, nccols),
data_format: if relaxed {
DataFormat::EncodeGappedPeakRs
} else {
DataFormat::EncodeGappedPeak
},
compliant_columns,
non_compliant_columns: nccols,
});
}
}
return Ok(BedClassificationOutput {
bed_compliance: format!("bed{}+{}", compliant_columns, nccols),
data_format: if relaxed {
if nccols == 0 {
DataFormat::UcscBedRs
} else {
DataFormat::BedLikeRs
}
} else {
DataFormat::BedLike
},
compliant_columns,
non_compliant_columns: nccols,
});
}
}
Ok(BedClassificationOutput {
bed_compliance: format!("bed{}+0", compliant_columns),
data_format: if relaxed {
DataFormat::UcscBedRs
} else {
DataFormat::UcscBed
},
compliant_columns,
non_compliant_columns: 0,
})
}
#[cfg(test)]
#[cfg(feature = "bedclassifier")]
mod tests {
use super::*;
fn get_test_path(file_name: &str) -> std::path::PathBuf {
std::env::current_dir()
.unwrap()
.join("../tests/data/regionset")
.join(file_name)
}
#[cfg(feature = "bedclassifier")]
#[test]
fn test_classify_bed_narrowpeak() {
let file_path = get_test_path("dummy.narrowPeak");
let region_set = RegionSet::try_from(file_path.to_str().unwrap()).unwrap();
let classification = classify_bed(®ion_set).unwrap();
println!("Classification: {}", classification);
println!("Bed compliance: {}", classification.bed_compliance);
println!("Data format: {}", classification.data_format);
println!("Compliant columns: {}", classification.compliant_columns);
println!(
"Non-compliant columns: {}",
classification.non_compliant_columns
);
assert!(classification.bed_compliance.starts_with("bed"));
assert_eq!(classification.data_format, DataFormat::EncodeNarrowPeak);
}
#[cfg(feature = "bedclassifier")]
#[test]
fn test_classify_bed_basic() {
let file_path = get_test_path("dummy_headers.bed");
let region_set = RegionSet::try_from(file_path.to_str().unwrap()).unwrap();
let classification = classify_bed(®ion_set).unwrap();
println!("Classification: {}", classification);
println!("Bed compliance: {}", classification.bed_compliance);
println!("Data format: {}", classification.data_format);
assert!(classification.bed_compliance.starts_with("bed"));
assert!(classification.compliant_columns >= 3); assert_eq!(classification.data_format, DataFormat::UcscBed);
}
}