use arrow::array::Array;
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
use std::fs::File;
use std::path::Path;
use crate::LoadError;
use crate::parquet_helpers::{
extract_optional_float64, extract_optional_int32, extract_required_int32,
};
#[derive(Debug, Clone, PartialEq)]
pub struct GenericConstraintBoundsRow {
pub constraint_id: i32,
pub stage_id: i32,
pub block_id: Option<i32>,
pub bound_lower: Option<f64>,
pub bound_upper: Option<f64>,
}
pub fn parse_generic_constraint_bounds(
path: &Path,
) -> Result<Vec<GenericConstraintBoundsRow>, LoadError> {
let file = File::open(path).map_err(|e| LoadError::io(path, e))?;
let builder = ParquetRecordBatchReaderBuilder::try_new(file)
.map_err(|e| LoadError::parse(path, e.to_string()))?;
let reader = builder
.build()
.map_err(|e| LoadError::parse(path, e.to_string()))?;
let mut rows: Vec<GenericConstraintBoundsRow> = Vec::new();
for batch_result in reader {
let batch = batch_result.map_err(|e| LoadError::parse(path, e.to_string()))?;
let constraint_id_col = extract_required_int32(&batch, "constraint_id", path)?;
let stage_id_col = extract_required_int32(&batch, "stage_id", path)?;
let block_id_col = extract_optional_int32(&batch, "block_id", path)?;
let bound_lower_col = extract_optional_float64(&batch, "bound_lower", path)?;
let bound_upper_col = extract_optional_float64(&batch, "bound_upper", path)?;
let n = batch.num_rows();
let base_idx = rows.len();
rows.reserve(n);
for i in 0..n {
let row_idx = base_idx + i;
let constraint_id = constraint_id_col.value(i);
let stage_id = stage_id_col.value(i);
let block_id = block_id_col
.filter(|col| !col.is_null(i))
.map(|col| col.value(i));
let bound_lower = bound_lower_col
.filter(|col| !col.is_null(i))
.map(|col| col.value(i));
let bound_upper = bound_upper_col
.filter(|col| !col.is_null(i))
.map(|col| col.value(i));
if let Some(bl) = bound_lower
&& !bl.is_finite()
{
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("generic_constraint_bounds[{row_idx}].bound_lower"),
message: format!("value must be finite, got {bl}"),
});
}
if let Some(bu) = bound_upper
&& !bu.is_finite()
{
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("generic_constraint_bounds[{row_idx}].bound_upper"),
message: format!("value must be finite, got {bu}"),
});
}
rows.push(GenericConstraintBoundsRow {
constraint_id,
stage_id,
block_id,
bound_lower,
bound_upper,
});
}
}
rows.sort_by(|a, b| {
a.constraint_id
.cmp(&b.constraint_id)
.then_with(|| a.stage_id.cmp(&b.stage_id))
.then_with(|| match (a.block_id, b.block_id) {
(None, None) => std::cmp::Ordering::Equal,
(None, Some(_)) => std::cmp::Ordering::Less,
(Some(_), None) => std::cmp::Ordering::Greater,
(Some(a), Some(b)) => a.cmp(&b),
})
});
Ok(rows)
}
#[cfg(test)]
#[allow(
clippy::doc_markdown,
clippy::expect_used,
clippy::panic,
clippy::too_many_lines,
clippy::unwrap_used
)]
mod tests {
use super::*;
use arrow::array::{Float64Array, Int32Array};
use arrow::datatypes::{DataType, Field, Schema};
use arrow::record_batch::RecordBatch;
use parquet::arrow::ArrowWriter;
use std::sync::Arc;
use tempfile::NamedTempFile;
fn schema_full() -> Arc<Schema> {
Arc::new(Schema::new(vec![
Field::new("constraint_id", DataType::Int32, false),
Field::new("stage_id", DataType::Int32, false),
Field::new("block_id", DataType::Int32, true), Field::new("bound_lower", DataType::Float64, true), Field::new("bound_upper", DataType::Float64, true), ]))
}
fn schema_no_endpoints() -> Arc<Schema> {
Arc::new(Schema::new(vec![
Field::new("constraint_id", DataType::Int32, false),
Field::new("stage_id", DataType::Int32, false),
]))
}
fn write_parquet(batch: &RecordBatch) -> NamedTempFile {
let tmp = NamedTempFile::new().expect("tempfile");
let mut writer = ArrowWriter::try_new(tmp.reopen().expect("reopen"), batch.schema(), None)
.expect("ArrowWriter");
writer.write(batch).expect("write batch");
writer.close().expect("close writer");
tmp
}
fn make_batch(
constraint_ids: &[i32],
stage_ids: &[i32],
block_ids: Vec<Option<i32>>,
bound_lowers: Vec<Option<f64>>,
bound_uppers: Vec<Option<f64>>,
) -> RecordBatch {
let block_arr: Int32Array = block_ids.into_iter().collect();
let lower_arr: Float64Array = bound_lowers.into_iter().collect();
let upper_arr: Float64Array = bound_uppers.into_iter().collect();
RecordBatch::try_new(
schema_full(),
vec![
Arc::new(Int32Array::from(constraint_ids.to_vec())),
Arc::new(Int32Array::from(stage_ids.to_vec())),
Arc::new(block_arr),
Arc::new(lower_arr),
Arc::new(upper_arr),
],
)
.expect("valid batch")
}
#[test]
fn test_parse_with_nullable_block_id() {
let batch = make_batch(
&[0, 0, 1],
&[0, 0, 2],
vec![Some(0), None, Some(1)],
vec![Some(100.0), Some(200.0), Some(300.0)],
vec![None, None, None],
);
let tmp = write_parquet(&batch);
let rows = parse_generic_constraint_bounds(tmp.path()).unwrap();
assert_eq!(rows.len(), 3);
assert_eq!(rows[0].constraint_id, 0);
assert_eq!(rows[0].stage_id, 0);
assert_eq!(rows[0].block_id, None);
assert_eq!(rows[0].bound_lower, Some(200.0));
assert_eq!(rows[1].constraint_id, 0);
assert_eq!(rows[1].stage_id, 0);
assert_eq!(rows[1].block_id, Some(0));
assert_eq!(rows[1].bound_lower, Some(100.0));
assert_eq!(rows[2].constraint_id, 1);
assert_eq!(rows[2].stage_id, 2);
assert_eq!(rows[2].block_id, Some(1));
assert_eq!(rows[2].bound_lower, Some(300.0));
}
#[test]
fn test_parse_without_block_id_column() {
let schema = Arc::new(Schema::new(vec![
Field::new("constraint_id", DataType::Int32, false),
Field::new("stage_id", DataType::Int32, false),
Field::new("bound_lower", DataType::Float64, true),
]));
let batch = RecordBatch::try_new(
schema,
vec![
Arc::new(Int32Array::from(vec![0, 1])),
Arc::new(Int32Array::from(vec![3, 4])),
Arc::new(Float64Array::from(vec![50.0, 75.0])),
],
)
.expect("valid batch");
let tmp = write_parquet(&batch);
let rows = parse_generic_constraint_bounds(tmp.path()).unwrap();
assert_eq!(rows.len(), 2);
assert!(rows[0].block_id.is_none());
assert!(rows[1].block_id.is_none());
}
#[test]
fn test_parse_non_finite_bound_lower_returns_schema_error() {
let batch = make_batch(&[0], &[0], vec![None], vec![Some(f64::NAN)], vec![None]);
let tmp = write_parquet(&batch);
let err = parse_generic_constraint_bounds(tmp.path()).unwrap_err();
match &err {
LoadError::SchemaError { field, message, .. } => {
assert_eq!(field, "generic_constraint_bounds[0].bound_lower");
assert!(
message.contains("finite"),
"message should mention 'finite', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn test_parse_infinite_bound_lower_returns_schema_error() {
let batch = make_batch(
&[0],
&[0],
vec![None],
vec![Some(f64::INFINITY)],
vec![None],
);
let tmp = write_parquet(&batch);
let err = parse_generic_constraint_bounds(tmp.path()).unwrap_err();
match &err {
LoadError::SchemaError { field, .. } => {
assert_eq!(field, "generic_constraint_bounds[0].bound_lower");
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn test_parse_non_finite_bound_upper_returns_schema_error() {
let batch = make_batch(
&[0],
&[0],
vec![None],
vec![Some(100.0)],
vec![Some(f64::NAN)],
);
let tmp = write_parquet(&batch);
let err = parse_generic_constraint_bounds(tmp.path()).unwrap_err();
match &err {
LoadError::SchemaError { field, message, .. } => {
assert_eq!(field, "generic_constraint_bounds[0].bound_upper");
assert!(
message.contains("finite"),
"message should mention 'finite', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn test_parse_missing_endpoint_columns_yields_none_endpoints() {
let batch = RecordBatch::try_new(
schema_no_endpoints(),
vec![
Arc::new(Int32Array::from(vec![0])),
Arc::new(Int32Array::from(vec![0])),
],
)
.expect("valid batch");
let tmp = write_parquet(&batch);
let rows = parse_generic_constraint_bounds(tmp.path()).unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].bound_lower, None);
assert_eq!(rows[0].bound_upper, None);
}
#[test]
fn test_parse_empty_parquet() {
let batch = make_batch(&[], &[], vec![], vec![], vec![]);
let tmp = write_parquet(&batch);
let rows = parse_generic_constraint_bounds(tmp.path()).unwrap();
assert!(rows.is_empty());
}
#[test]
fn test_parse_sort_order_invariance() {
let batch = make_batch(
&[2, 0, 1, 0],
&[0, 1, 0, 0],
vec![None, Some(0), None, None],
vec![Some(10.0), Some(20.0), Some(30.0), Some(40.0)],
vec![None, None, None, None],
);
let tmp = write_parquet(&batch);
let rows = parse_generic_constraint_bounds(tmp.path()).unwrap();
assert_eq!(rows.len(), 4);
assert_eq!(
(rows[0].constraint_id, rows[0].stage_id, rows[0].block_id),
(0, 0, None)
);
assert_eq!(
(rows[1].constraint_id, rows[1].stage_id, rows[1].block_id),
(0, 1, Some(0))
);
assert_eq!(
(rows[2].constraint_id, rows[2].stage_id, rows[2].block_id),
(1, 0, None)
);
assert_eq!(
(rows[3].constraint_id, rows[3].stage_id, rows[3].block_id),
(2, 0, None)
);
}
#[test]
fn test_parse_lower_only_row() {
let batch = make_batch(&[0], &[0], vec![None], vec![Some(5.0)], vec![None]);
let tmp = write_parquet(&batch);
let rows = parse_generic_constraint_bounds(tmp.path()).unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].bound_lower, Some(5.0));
assert_eq!(rows[0].bound_upper, None);
}
#[test]
fn test_parse_upper_only_row() {
let batch = make_batch(&[0], &[0], vec![None], vec![None], vec![Some(10.0)]);
let tmp = write_parquet(&batch);
let rows = parse_generic_constraint_bounds(tmp.path()).unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].bound_lower, None);
assert_eq!(rows[0].bound_upper, Some(10.0));
}
#[test]
fn test_parse_two_sided_row() {
let batch = make_batch(&[0], &[0], vec![None], vec![Some(35.0)], vec![Some(100.0)]);
let tmp = write_parquet(&batch);
let rows = parse_generic_constraint_bounds(tmp.path()).unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].bound_lower, Some(35.0));
assert_eq!(rows[0].bound_upper, Some(100.0));
}
#[test]
fn test_parse_both_endpoints_null_row() {
let batch = make_batch(&[0], &[0], vec![None], vec![None], vec![None]);
let tmp = write_parquet(&batch);
let rows = parse_generic_constraint_bounds(tmp.path()).unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].bound_lower, None);
assert_eq!(rows[0].bound_upper, None);
}
#[test]
fn test_parse_sort_order_preserved_with_both_endpoints_present() {
let batch = make_batch(
&[2, 0, 1, 0],
&[0, 1, 0, 0],
vec![None, Some(0), None, None],
vec![Some(10.0), Some(20.0), Some(30.0), Some(40.0)],
vec![Some(15.0), None, Some(35.0), None],
);
let tmp = write_parquet(&batch);
let rows = parse_generic_constraint_bounds(tmp.path()).unwrap();
assert_eq!(rows.len(), 4);
assert_eq!(
(rows[0].constraint_id, rows[0].stage_id, rows[0].block_id),
(0, 0, None)
);
assert_eq!(rows[0].bound_upper, None);
assert_eq!(
(rows[1].constraint_id, rows[1].stage_id, rows[1].block_id),
(0, 1, Some(0))
);
assert_eq!(rows[1].bound_upper, None);
assert_eq!(
(rows[2].constraint_id, rows[2].stage_id, rows[2].block_id),
(1, 0, None)
);
assert_eq!(rows[2].bound_upper, Some(35.0));
assert_eq!(
(rows[3].constraint_id, rows[3].stage_id, rows[3].block_id),
(2, 0, None)
);
assert_eq!(rows[3].bound_upper, Some(15.0));
}
}