use crate::dataframe::DataFrame;
use crate::error::{Error, Result};
use crate::series::Series;
use std::collections::HashMap;
use std::path::Path;
use super::core::read_parquet;
#[derive(Debug, Clone)]
pub enum PredicateFilter {
Equals(String, String),
Range(String, String, String),
In(String, Vec<String>),
NotNull(String),
Custom(String),
}
#[derive(Debug, Clone)]
pub struct SchemaEvolution {
pub source_schema: String,
pub target_schema: String,
pub column_mappings: HashMap<String, String>,
pub columns_to_add: HashMap<String, String>,
pub columns_to_remove: Vec<String>,
pub type_conversions: HashMap<String, String>,
}
impl Default for SchemaEvolution {
fn default() -> Self {
Self {
source_schema: String::new(),
target_schema: String::new(),
column_mappings: HashMap::new(),
columns_to_add: HashMap::new(),
columns_to_remove: Vec::new(),
type_conversions: HashMap::new(),
}
}
}
pub fn read_parquet_with_schema_evolution(
path: impl AsRef<Path>,
schema_evolution: SchemaEvolution,
) -> Result<DataFrame> {
let mut df = read_parquet(path.as_ref())?;
apply_schema_evolution(&mut df, &schema_evolution)?;
Ok(df)
}
pub(super) fn apply_schema_evolution(
df: &mut DataFrame,
evolution: &SchemaEvolution,
) -> Result<()> {
for (old_name, new_name) in &evolution.column_mappings {
if old_name == new_name {
continue;
}
if df.contains_column(old_name) {
let mut mapping = HashMap::new();
mapping.insert(old_name.clone(), new_name.clone());
df.rename_columns(&mapping)?;
}
}
if !evolution.columns_to_remove.is_empty() {
*df = drop_columns_preserving_type(df, &evolution.columns_to_remove)?;
}
for (col_name, default_value) in &evolution.columns_to_add {
let row_count = df.row_count();
let default_values = vec![default_value.clone(); row_count];
let series = Series::new(default_values, Some(col_name.clone()))?;
df.add_column(col_name.clone(), series)?;
}
Ok(())
}
fn drop_columns_preserving_type(df: &DataFrame, names: &[String]) -> Result<DataFrame> {
let mut result = DataFrame::new();
for col_name in df.column_names() {
if names.iter().any(|n| n == col_name) {
continue;
}
if let Ok(series) = df.get_column::<i64>(col_name) {
result.add_column(col_name.clone(), series.clone())?;
} else if let Ok(series) = df.get_column::<f64>(col_name) {
result.add_column(col_name.clone(), series.clone())?;
} else if let Ok(series) = df.get_column::<bool>(col_name) {
result.add_column(col_name.clone(), series.clone())?;
} else if let Ok(series) = df.get_column::<String>(col_name) {
result.add_column(col_name.clone(), series.clone())?;
} else {
return Err(Error::Type(format!(
"Column '{}' has a type not supported by schema-evolution column removal",
col_name
)));
}
}
Ok(result)
}
pub fn read_parquet_with_predicates(
path: impl AsRef<Path>,
predicates: Vec<PredicateFilter>,
) -> Result<DataFrame> {
let df = read_parquet(path.as_ref())?;
apply_predicate_filters(df, &predicates)
}
fn predicate_cell_equals(cell: &str, target: &str) -> bool {
if cell == target {
return true;
}
matches!(
(cell.trim().parse::<f64>(), target.trim().parse::<f64>()),
(Ok(a), Ok(b)) if a == b
)
}
pub(super) fn apply_predicate_filters(
df: DataFrame,
predicates: &[PredicateFilter],
) -> Result<DataFrame> {
if predicates.is_empty() {
return Ok(df);
}
let row_count = df.row_count();
let mut keep = vec![true; row_count];
for predicate in predicates {
match predicate {
PredicateFilter::Equals(column, value) => {
let values = df.get_column_string_values(column)?;
for (i, cell) in values.iter().enumerate() {
if !predicate_cell_equals(cell, value) {
keep[i] = false;
}
}
}
PredicateFilter::Range(column, min, max) => {
let values = df.get_column_string_values(column)?;
let min_f = min.trim().parse::<f64>().ok();
let max_f = max.trim().parse::<f64>().ok();
for (i, cell) in values.iter().enumerate() {
let in_range = match cell.trim().parse::<f64>() {
Ok(v) => {
min_f.map_or(true, |lo| v >= lo) && max_f.map_or(true, |hi| v <= hi)
}
Err(_) => cell.as_str() >= min.as_str() && cell.as_str() <= max.as_str(),
};
if !in_range {
keep[i] = false;
}
}
}
PredicateFilter::In(column, allowed) => {
let values = df.get_column_string_values(column)?;
for (i, cell) in values.iter().enumerate() {
if !allowed.iter().any(|a| predicate_cell_equals(cell, a)) {
keep[i] = false;
}
}
}
PredicateFilter::NotNull(column) => {
let values = df.get_column_string_values(column)?;
for (i, cell) in values.iter().enumerate() {
let trimmed = cell.trim();
if trimmed.is_empty()
|| trimmed.eq_ignore_ascii_case("null")
|| trimmed.eq_ignore_ascii_case("nan")
{
keep[i] = false;
}
}
}
PredicateFilter::Custom(expression) => {
return Err(Error::NotImplemented(format!(
"Custom predicate expression '{}' is not supported",
expression
)));
}
}
}
let kept_indices: Vec<usize> = (0..row_count).filter(|&i| keep[i]).collect();
df.sample(&kept_indices)
}