use csv::{ReaderBuilder, StringRecord};
use ndarray::{Array2, ArrayViewMut1, Axis, s};
use rayon::prelude::*;
use serde::{Deserialize, Serialize};
use std::cmp::Ordering;
use std::collections::{HashMap, HashSet};
use std::fmt;
use std::path::Path;
fn natural_level_cmp(a: &str, b: &str) -> Ordering {
let mut ia = 0;
let mut ib = 0;
let ba = a.as_bytes();
let bb = b.as_bytes();
while ia < ba.len() && ib < bb.len() {
if ba[ia].is_ascii_digit() && bb[ib].is_ascii_digit() {
let sa = ia;
let sb = ib;
while ia < ba.len() && ba[ia].is_ascii_digit() {
ia += 1;
}
while ib < bb.len() && bb[ib].is_ascii_digit() {
ib += 1;
}
let da = &a[sa..ia];
let db = &b[sb..ib];
let ta = da.trim_start_matches('0');
let tb = db.trim_start_matches('0');
let ta = if ta.is_empty() { "0" } else { ta };
let tb = if tb.is_empty() { "0" } else { tb };
match ta.len().cmp(&tb.len()).then_with(|| ta.cmp(tb)) {
Ordering::Equal if da.len() != db.len() => return da.len().cmp(&db.len()),
Ordering::Equal => {}
ord => return ord,
}
} else {
match ba[ia].cmp(&bb[ib]) {
Ordering::Equal => {
ia += 1;
ib += 1;
}
ord => return ord,
}
}
}
ba.len().cmp(&bb.len())
}
fn sort_levels_canonical(levels: &mut [String]) {
levels.sort_by(|a, b| natural_level_cmp(a, b));
}
pub fn encode_optional_categorical_column(
name: &str,
column: &[Option<&str>],
) -> Result<(SchemaColumn, Vec<f64>), DataError> {
if column.is_empty() {
return Err(DataError::EmptyInput {
reason: "table data cannot be empty".to_string(),
});
}
let mut levels = Vec::new();
for (row, label) in column.iter().enumerate() {
let Some(label) = label else {
continue;
};
let label = label.trim();
if label.is_empty() {
return Err(DataError::EmptyInput {
reason: format!("empty field at row {}, column '{name}'", row + 1),
});
}
levels.push(label.to_string());
}
sort_levels_canonical(&mut levels);
levels.dedup();
let level_map = levels
.iter()
.enumerate()
.map(|(index, level)| (level.as_str(), index as f64))
.collect::<HashMap<_, _>>();
let values = column
.iter()
.map(|value| match value {
None => Ok(f64::NAN),
Some(label) => level_map.get(label.trim()).copied().ok_or_else(|| {
DataError::EncodingFailure {
reason: format!(
"internal: level '{}' missing from freshly built map for column '{name}'",
label.trim()
),
}
}),
})
.collect::<Result<Vec<_>, _>>()?;
Ok((
SchemaColumn {
name: name.to_string(),
kind: ColumnKindTag::Categorical,
levels,
},
values,
))
}
#[inline]
pub fn canonical_level_bits(v: f64) -> u64 {
if v == 0.0 {
0.0_f64.to_bits()
} else if v.is_nan() {
f64::NAN.to_bits()
} else {
v.to_bits()
}
}
#[derive(Debug, Clone)]
pub enum DataError {
SchemaMismatch { reason: String },
ParseError { reason: String },
EncodingFailure { reason: String },
EmptyInput { reason: String },
InvalidValue { reason: String },
DegenerateColumn { column: String, problem: String },
ColumnNotFound {
name: String,
role: Option<String>,
available: Vec<String>,
similar: Vec<String>,
tsv_hint: bool,
},
}
impl DataError {
#[must_use]
fn with_source_path(self, path: &Path) -> Self {
let qualify = |reason: String| {
if reason.contains(&path.display().to_string()) {
reason
} else {
format!("data file '{}': {reason}", path.display())
}
};
match self {
Self::SchemaMismatch { reason } => Self::SchemaMismatch { reason: qualify(reason) },
Self::ParseError { reason } => Self::ParseError { reason: qualify(reason) },
Self::EncodingFailure { reason } => Self::EncodingFailure { reason: qualify(reason) },
Self::EmptyInput { reason } => Self::EmptyInput { reason: qualify(reason) },
Self::InvalidValue { reason } => Self::InvalidValue { reason: qualify(reason) },
column @ Self::ColumnNotFound { .. } => column,
degenerate @ Self::DegenerateColumn { .. } => degenerate,
}
}
#[must_use]
pub fn advice(&self) -> Option<String> {
match self {
Self::SchemaMismatch { .. } => Some(
"Verify the new data has the same columns and types as the training data \
and that the formula terms match."
.to_string(),
),
Self::ParseError { .. }
| Self::EncodingFailure { .. }
| Self::EmptyInput { .. }
| Self::InvalidValue { .. }
| Self::DegenerateColumn { .. }
| Self::ColumnNotFound { .. } => None,
}
}
pub fn column_not_found(
col_map: &HashMap<String, usize>,
name: &str,
role: Option<&str>,
) -> Self {
let target_lower = name.to_lowercase();
let mut similar: Vec<String> = col_map
.keys()
.filter(|k| {
let k_lower = k.to_lowercase();
k_lower.contains(&target_lower)
|| target_lower.contains(&k_lower)
|| shared_prefix(&k_lower, &target_lower) >= 3
})
.cloned()
.collect();
similar.sort_unstable();
let mut available: Vec<String> = col_map.keys().cloned().collect();
available.sort_unstable();
let tsv_hint = available.len() == 1 && available[0].contains('\t');
Self::ColumnNotFound {
name: name.to_string(),
role: role.map(str::to_string),
available,
similar,
tsv_hint,
}
}
}
impl fmt::Display for DataError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
DataError::SchemaMismatch { reason }
| DataError::ParseError { reason }
| DataError::EncodingFailure { reason }
| DataError::EmptyInput { reason }
| DataError::InvalidValue { reason } => f.write_str(reason),
DataError::DegenerateColumn { column, problem } => {
write!(f, "column '{column}' {problem}")
}
DataError::ColumnNotFound {
name,
role,
available,
similar,
tsv_hint,
} => {
let label = match role {
Some(r) => format!("{r} column '{name}'"),
None => format!("column '{name}'"),
};
let tsv_suffix = if *tsv_hint {
" — your file appears to be tab-separated; gam expects comma-separated CSV. \
Replace tabs with commas, or pre-convert with `tr '\\t' ',' < file.tsv > file.csv`."
} else {
""
};
if similar.is_empty() {
write!(
f,
"{label} not found in data. Available columns: [{}]{tsv_suffix}",
available.join(", ")
)
} else {
write!(
f,
"{label} not found in data. Did you mean one of [{}]? Full list: [{}]{tsv_suffix}",
similar.join(", "),
available.join(", ")
)
}
}
}
}
}
impl std::error::Error for DataError {}
impl From<DataError> for String {
fn from(err: DataError) -> String {
err.to_string()
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct DataSchema {
pub columns: Vec<SchemaColumn>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct SchemaColumn {
pub name: String,
pub kind: ColumnKindTag,
#[serde(default)]
pub levels: Vec<String>,
}
#[derive(Clone, Copy, Debug, Serialize, Deserialize, Eq, PartialEq)]
#[serde(rename_all = "kebab-case")]
pub enum ColumnKindTag {
Continuous,
Binary,
Categorical,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum UnseenCategoryPolicy {
Error,
EncodeUnknownForColumns(HashSet<String>),
}
impl UnseenCategoryPolicy {
pub fn encode_unknown_for_columns(columns: HashSet<String>) -> Self {
if columns.is_empty() {
Self::Error
} else {
Self::EncodeUnknownForColumns(columns)
}
}
fn unseen_code_for(&self, column_name: &str, level_count: usize) -> Option<f64> {
match self {
Self::Error => None,
Self::EncodeUnknownForColumns(columns) => {
columns.contains(column_name).then_some(level_count as f64)
}
}
}
}
#[derive(Clone, Debug)]
pub struct EncodedDataset {
pub headers: Vec<String>,
pub values: Array2<f64>,
pub schema: DataSchema,
pub column_kinds: Vec<ColumnKindTag>,
}
impl EncodedDataset {
pub fn validate_fit_boundary(&self) -> Result<(), DataError> {
if self.headers.is_empty() {
return Err(DataError::DegenerateColumn {
column: "<table>".to_string(),
problem: "has no columns".to_string(),
});
}
let mut seen = HashSet::with_capacity(self.headers.len());
for name in &self.headers {
if !seen.insert(name.as_str()) {
return Err(DataError::DegenerateColumn {
column: name.clone(),
problem: "has a duplicate name".to_string(),
});
}
}
if self.values.nrows() == 0 {
return Err(DataError::DegenerateColumn {
column: "<table>".to_string(),
problem: "has no observations".to_string(),
});
}
if self.values.ncols() != self.headers.len() {
return Err(DataError::SchemaMismatch {
reason: format!(
"table has {} headers but {} value columns",
self.headers.len(),
self.values.ncols()
),
});
}
for (index, name) in self.headers.iter().enumerate() {
if self.column_kinds.get(index) == Some(&ColumnKindTag::Categorical)
&& self
.schema
.columns
.get(index)
.is_some_and(|column| column.levels.len() < 2)
{
return Err(DataError::DegenerateColumn {
column: name.clone(),
problem: "is a factor with fewer than two levels".to_string(),
});
}
let column = self.values.column(index);
let finite_count = column.iter().filter(|value| value.is_finite()).count();
if finite_count == 1 && column.len() > 1 {
return Err(DataError::DegenerateColumn {
column: name.clone(),
problem: "has only one non-missing value".to_string(),
});
}
if let Some((row, value)) = column
.iter()
.enumerate()
.find(|(_, value)| !value.is_finite())
{
return Err(DataError::DegenerateColumn {
column: name.clone(),
problem: format!("has non-finite value {value} at row {}", row + 1),
});
}
}
Ok(())
}
pub fn column_map(&self) -> HashMap<String, usize> {
self.headers
.iter()
.enumerate()
.map(|(index, header)| (header.clone(), index))
.collect()
}
pub fn feature_ranges(&self) -> Vec<(f64, f64)> {
self.values
.axis_iter(Axis(1))
.into_par_iter()
.map(|col| {
let (lo, hi) =
col.iter()
.fold((f64::INFINITY, f64::NEG_INFINITY), |(lo, hi), &v| {
if v.is_finite() {
(lo.min(v), hi.max(v))
} else {
(lo, hi)
}
});
if !lo.is_finite() || !hi.is_finite() {
(0.0, 0.0)
} else {
(lo, hi)
}
})
.collect()
}
}
fn shared_prefix(a: &str, b: &str) -> usize {
a.chars()
.zip(b.chars())
.take_while(|(ca, cb)| ca == cb)
.count()
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum DataFormat {
Csv,
Tsv,
Parquet,
}
fn detect_format(path: &Path) -> Result<DataFormat, DataError> {
let ext = path
.extension()
.and_then(|s| s.to_str())
.unwrap_or_default()
.to_ascii_lowercase();
match ext.as_str() {
"csv" => Ok(DataFormat::Csv),
"tsv" | "txt" | "tab" => Ok(DataFormat::Tsv),
"parquet" | "pq" | "pqt" => Ok(DataFormat::Parquet),
other => Err(DataError::ParseError {
reason: format!(
"unsupported data file extension '.{other}'; expected csv, tsv, txt, parquet, or pq: '{}'",
path.display()
),
}),
}
}
pub fn load_dataset_projected(
path: &Path,
requested_columns: &[String],
) -> Result<EncodedDataset, DataError> {
load_dataset_projected_with_categorical_roles(path, requested_columns, &HashSet::new())
}
pub fn load_dataset_projected_with_categorical_roles(
path: &Path,
requested_columns: &[String],
categorical_roles: &HashSet<&str>,
) -> Result<EncodedDataset, DataError> {
(match detect_format(path)? {
DataFormat::Csv => {
load_delimited_inferred(path, b',', requested_columns, categorical_roles)
}
DataFormat::Tsv => {
load_delimited_inferred(path, b'\t', requested_columns, categorical_roles)
}
DataFormat::Parquet => load_parquet_inferred(path, requested_columns, categorical_roles),
})
.map_err(|error| error.with_source_path(path))
}
pub fn load_datasetwith_schema_projected(
path: &Path,
schema: &DataSchema,
unseen_policy: UnseenCategoryPolicy,
requested_columns: &[String],
) -> Result<EncodedDataset, DataError> {
(match detect_format(path)? {
DataFormat::Csv => {
load_delimited_with_schema(path, b',', schema, unseen_policy, requested_columns)
}
DataFormat::Tsv => {
load_delimited_with_schema(path, b'\t', schema, unseen_policy, requested_columns)
}
DataFormat::Parquet => {
load_parquet_with_schema(path, schema, unseen_policy, requested_columns)
}
})
.map_err(|error| error.with_source_path(path))
}
pub fn load_csvwith_inferred_schema(path: &Path) -> Result<EncodedDataset, DataError> {
load_delimited_inferred(path, b',', &[], &HashSet::new())
.map_err(|error| error.with_source_path(path))
}
pub const CATEGORICAL_CELL_SENTINEL: char = '\u{0}';
pub fn strip_categorical_sentinel(cell: &str) -> (&str, bool) {
match cell.strip_prefix(CATEGORICAL_CELL_SENTINEL) {
Some(rest) => (rest, true),
None => (cell, false),
}
}
fn resolve_requested_columns(
all_headers: &[String],
requested_columns: &[String],
) -> Result<Vec<usize>, DataError> {
if requested_columns.is_empty() {
return Ok((0..all_headers.len()).collect());
}
let requested_set: HashSet<&str> = requested_columns.iter().map(String::as_str).collect();
let mut selected = Vec::with_capacity(requested_set.len());
for (idx, name) in all_headers.iter().enumerate() {
if requested_set.contains(name.as_str()) {
selected.push(idx);
}
}
if selected.len() != requested_set.len() {
let available_map: HashMap<String, usize> = all_headers
.iter()
.enumerate()
.map(|(index, header)| (header.clone(), index))
.collect();
let missing = requested_columns
.iter()
.filter(|name| !available_map.contains_key(name.as_str()))
.map(|name| {
DataError::column_not_found(&available_map, name, Some("requested")).to_string()
})
.collect::<Vec<_>>();
return Err(DataError::SchemaMismatch {
reason: missing.join("; "),
});
}
Ok(selected)
}
fn projected_headers(all_headers: &[String], selected_indices: &[usize]) -> Vec<String> {
selected_indices
.iter()
.map(|&idx| all_headers[idx].clone())
.collect()
}
fn load_delimited_inferred(
path: &Path,
delimiter: u8,
requested_columns: &[String],
categorical_roles: &HashSet<&str>,
) -> Result<EncodedDataset, DataError> {
let t_open = std::time::Instant::now();
let mut rdr = ReaderBuilder::new()
.has_headers(true)
.delimiter(delimiter)
.from_path(path)
.map_err(|e| DataError::ParseError {
reason: format!("failed to open '{}': {e}", path.display()),
})?;
let all_headers: Vec<String> = rdr
.headers()
.map_err(|e| DataError::ParseError {
reason: format!("failed to read headers: {e}"),
})?
.iter()
.map(|s| s.trim().to_string())
.collect();
if all_headers.is_empty() {
return Err(DataError::EmptyInput {
reason: "file has no headers".to_string(),
});
}
let selected_indices = resolve_requested_columns(&all_headers, requested_columns)?;
let headers = projected_headers(&all_headers, &selected_indices);
let p = headers.len();
let open_ms = t_open.elapsed().as_secs_f64() * 1000.0;
if open_ms > 100.0 {
log::info!(
"[DATA-LOAD] delim_open+headers | n_headers={} | n_proj={} | {:.1}ms",
all_headers.len(),
p,
open_ms
);
}
let mut inference = vec![DelimitedInferenceState::default(); p];
let mut total_rows: usize = 0;
let t_stream = std::time::Instant::now();
let mut record = StringRecord::new();
while rdr
.read_record(&mut record)
.map_err(|e| DataError::ParseError {
reason: format!("failed reading row: {e}"),
})?
{
if record.len() != all_headers.len() {
return Err(DataError::SchemaMismatch {
reason: format!(
"row width mismatch at row {}: got {} fields, expected {}",
total_rows + 1,
record.len(),
all_headers.len()
),
});
}
total_rows += 1;
for (j, &selected_idx) in selected_indices.iter().enumerate() {
inference[j].observe(
record
.get(selected_idx)
.expect("record width was checked against the header row above")
.trim(),
total_rows,
&headers[j],
)?;
}
}
let stream_ms = t_stream.elapsed().as_secs_f64() * 1000.0;
if stream_ms > 100.0 {
log::info!(
"[DATA-LOAD] delim_stream | n_rows={} | n_cols={} | {:.1}ms",
total_rows,
p,
stream_ms
);
}
if total_rows == 0 {
return Err(DataError::EmptyInput {
reason: "file has no rows".to_string(),
});
}
let t_schema = std::time::Instant::now();
let column_kinds = inference
.iter()
.enumerate()
.map(|(j, state)| state.kind(categorical_roles.contains(headers[j].as_str())))
.collect::<Vec<_>>();
let schema_ms = t_schema.elapsed().as_secs_f64() * 1000.0;
if schema_ms > 100.0 {
let n_cat = column_kinds
.iter()
.filter(|k| matches!(k, ColumnKindTag::Categorical))
.count();
log::info!(
"[DATA-LOAD] delim_convert+infer | n_cols={} | n_cat={} | {:.1}ms",
p,
n_cat,
schema_ms
);
}
let t_assemble = std::time::Instant::now();
let mut values = Array2::<f64>::zeros((total_rows, p));
let mut categorical_encoders = (0..p)
.map(|j| {
matches!(column_kinds[j], ColumnKindTag::Categorical).then(CategoricalEncoder::default)
})
.collect::<Vec<_>>();
let mut encode_rdr = ReaderBuilder::new()
.has_headers(true)
.delimiter(delimiter)
.from_path(path)
.map_err(|e| DataError::ParseError {
reason: format!("failed to reopen '{}': {e}", path.display()),
})?;
encode_rdr.headers().map_err(|e| DataError::ParseError {
reason: format!("failed to reread headers: {e}"),
})?;
let mut encoded_rows = 0usize;
while encode_rdr
.read_record(&mut record)
.map_err(|e| DataError::ParseError {
reason: format!("failed reading row: {e}"),
})?
{
if record.len() != all_headers.len() {
return Err(DataError::SchemaMismatch {
reason: format!(
"row width mismatch at row {}: got {} fields, expected {}",
encoded_rows + 1,
record.len(),
all_headers.len()
),
});
}
if encoded_rows >= total_rows {
return Err(DataError::SchemaMismatch {
reason: "data file changed while its schema was being discovered".to_string(),
});
}
for (j, &selected_idx) in selected_indices.iter().enumerate() {
let raw = record
.get(selected_idx)
.expect("record width was checked against the header row above")
.trim();
values[[encoded_rows, j]] = match column_kinds[j] {
ColumnKindTag::Continuous | ColumnKindTag::Binary => {
parse_inferred_numeric_cell(raw, encoded_rows + 1, &headers[j])?
}
ColumnKindTag::Categorical => {
if raw.is_empty() {
return Err(DataError::EmptyInput {
reason: format!(
"empty field at row {}, column '{}'",
encoded_rows + 1,
&headers[j]
),
});
}
categorical_encoders[j]
.as_mut()
.expect("categorical encoder")
.encode(raw) as f64
}
};
}
encoded_rows += 1;
}
if encoded_rows != total_rows {
return Err(DataError::SchemaMismatch {
reason: "data file changed while its schema was being discovered".to_string(),
});
}
let mut levels = vec![Vec::<String>::new(); p];
for (j, encoder) in categorical_encoders.into_iter().enumerate() {
if let Some(encoder) = encoder {
levels[j] = encoder.finish(values.column_mut(j), LevelOrder::Canonical);
}
}
let assemble_ms = t_assemble.elapsed().as_secs_f64() * 1000.0;
if assemble_ms > 100.0 {
log::info!(
"[DATA-LOAD] delim_assemble_array2 | n_rows={} | n_cols={} | {:.1}ms",
total_rows,
p,
assemble_ms
);
}
let schema = DataSchema {
columns: headers
.iter()
.enumerate()
.map(|(j, name)| SchemaColumn {
name: name.clone(),
kind: column_kinds[j],
levels: std::mem::take(&mut levels[j]),
})
.collect(),
};
Ok(EncodedDataset {
headers,
values,
schema,
column_kinds,
})
}
#[derive(Clone, Copy)]
struct DelimitedInferenceState {
all_numeric: bool,
all_binary: bool,
saw_numeric: bool,
}
impl Default for DelimitedInferenceState {
fn default() -> Self {
Self {
all_numeric: true,
all_binary: true,
saw_numeric: false,
}
}
}
impl DelimitedInferenceState {
fn observe(&mut self, raw: &str, row: usize, header: &str) -> Result<(), DataError> {
if raw.is_empty() {
return Err(DataError::EmptyInput {
reason: format!("empty field at row {row}, column '{header}'"),
});
}
if is_missing_marker(raw) {
return Ok(());
}
match raw.parse::<f64>() {
Ok(value) => {
self.saw_numeric = true;
if !value.is_finite() {
return Err(DataError::InvalidValue {
reason: format!("non-finite value at row {row}, column '{header}'"),
});
}
if (value - 0.0).abs() >= 1e-12 && (value - 1.0).abs() >= 1e-12 {
self.all_binary = false;
}
}
Err(_) => {
self.all_numeric = false;
self.all_binary = false;
}
}
Ok(())
}
fn kind(self, force_categorical: bool) -> ColumnKindTag {
if force_categorical || !(self.all_numeric && self.saw_numeric) {
ColumnKindTag::Categorical
} else if self.all_binary {
ColumnKindTag::Binary
} else {
ColumnKindTag::Continuous
}
}
}
fn is_missing_marker(raw: &str) -> bool {
matches!(
raw.trim().to_ascii_uppercase().as_str(),
"NA" | "N/A" | "NULL"
)
}
fn parse_inferred_numeric_cell(raw: &str, row: usize, header: &str) -> Result<f64, DataError> {
if raw.is_empty() {
return Err(DataError::EmptyInput {
reason: format!("empty field at row {row}, column '{header}'"),
});
}
if is_missing_marker(raw) {
return Ok(f64::NAN);
}
let value = raw
.parse::<f64>()
.map_err(|error| DataError::EncodingFailure {
reason: format!(
"failed to parse numeric value '{raw}' at row {row}, column '{header}': {error}"
),
})?;
if !value.is_finite() {
return Err(DataError::InvalidValue {
reason: format!("non-finite value at row {row}, column '{header}'"),
});
}
Ok(value)
}
#[derive(Clone, Copy)]
enum LevelOrder {
Encounter,
Canonical,
}
#[derive(Default)]
struct CategoricalEncoder {
encounter_codes: HashMap<String, usize>,
}
impl CategoricalEncoder {
fn encode(&mut self, label: &str) -> usize {
if let Some(&code) = self.encounter_codes.get(label) {
return code;
}
let code = self.encounter_codes.len();
self.encounter_codes.insert(label.to_owned(), code);
code
}
fn finish(self, mut encoded: ArrayViewMut1<'_, f64>, order: LevelOrder) -> Vec<String> {
match order {
LevelOrder::Encounter => {
let mut levels = std::iter::repeat_with(|| None)
.take(self.encounter_codes.len())
.collect::<Vec<Option<String>>>();
for (level, old_code) in self.encounter_codes {
levels[old_code] = Some(level);
}
levels
.into_iter()
.map(|level| level.expect("encounter code must name one level"))
.collect()
}
LevelOrder::Canonical => {
let mut levels_with_old_codes =
self.encounter_codes.into_iter().collect::<Vec<_>>();
levels_with_old_codes
.sort_by(|(a, _), (b, _)| natural_level_cmp(a.as_str(), b.as_str()));
let mut remap = vec![0usize; levels_with_old_codes.len()];
for (new_code, (_, old_code)) in levels_with_old_codes.iter().enumerate() {
remap[*old_code] = new_code;
}
for code in encoded.iter_mut() {
if code.is_finite() {
*code = remap[*code as usize] as f64;
}
}
levels_with_old_codes
.into_iter()
.map(|(level, _)| level)
.collect()
}
}
}
}
fn load_delimited_with_schema(
path: &Path,
delimiter: u8,
schema: &DataSchema,
unseen_policy: UnseenCategoryPolicy,
requested_columns: &[String],
) -> Result<EncodedDataset, DataError> {
let t_open = std::time::Instant::now();
let mut rdr = ReaderBuilder::new()
.has_headers(true)
.delimiter(delimiter)
.from_path(path)
.map_err(|e| DataError::ParseError {
reason: format!("failed to open '{}': {e}", path.display()),
})?;
let all_headers: Vec<String> = rdr
.headers()
.map_err(|e| DataError::ParseError {
reason: format!("failed to read headers: {e}"),
})?
.iter()
.map(|s| s.trim().to_string())
.collect();
if all_headers.is_empty() {
return Err(DataError::EmptyInput {
reason: "file has no headers".to_string(),
});
}
let selected_indices = resolve_requested_columns(&all_headers, requested_columns)?;
let headers = projected_headers(&all_headers, &selected_indices);
let p = headers.len();
let open_ms = t_open.elapsed().as_secs_f64() * 1000.0;
if open_ms > 100.0 {
log::info!(
"[DATA-LOAD] delim_schema_open+headers | n_headers={} | n_proj={} | {:.1}ms",
all_headers.len(),
p,
open_ms
);
}
let schema_byname: HashMap<&str, &SchemaColumn> = schema
.columns
.iter()
.map(|c| (c.name.as_str(), c))
.collect();
let mut col_meta = Vec::<ColMeta>::with_capacity(p);
for name in &headers {
if let Some(sc) = schema_byname.get(name.as_str()) {
let level_map = if matches!(sc.kind, ColumnKindTag::Categorical) {
Some(
sc.levels
.iter()
.enumerate()
.map(|(idx, v)| (v.as_str(), idx as f64))
.collect::<HashMap<_, _>>(),
)
} else {
None
};
col_meta.push(ColMeta {
kind: sc.kind,
level_map,
schema_col: (*sc).clone(),
});
} else {
col_meta.push(ColMeta {
kind: ColumnKindTag::Continuous, level_map: None,
schema_col: SchemaColumn {
name: name.clone(),
kind: ColumnKindTag::Continuous,
levels: Vec::new(),
},
});
}
}
let needs_inference: Vec<bool> = headers
.iter()
.map(|h| !schema_byname.contains_key(h.as_str()))
.collect();
if needs_inference.iter().all(|needs| !needs) {
let t_stream = std::time::Instant::now();
let mut flat_values = Vec::<f64>::new();
let mut total_rows = 0usize;
let mut record = StringRecord::new();
while rdr
.read_record(&mut record)
.map_err(|e| DataError::ParseError {
reason: format!("failed reading row: {e}"),
})?
{
if record.len() != all_headers.len() {
return Err(DataError::SchemaMismatch {
reason: format!(
"row width mismatch at row {}: got {} fields, expected {}",
total_rows + 1,
record.len(),
all_headers.len()
),
});
}
total_rows += 1;
for j in 0..p {
let raw = record
.get(selected_indices[j])
.expect("record width was checked against the header row above")
.trim();
flat_values.push(parse_cell_with_schema(
raw,
&col_meta[j],
total_rows,
&headers[j],
&unseen_policy,
)?);
}
}
if total_rows == 0 {
return Err(DataError::EmptyInput {
reason: "file has no rows".to_string(),
});
}
let values = Array2::from_shape_vec((total_rows, p), flat_values).map_err(|error| {
DataError::EncodingFailure {
reason: format!("failed to assemble schema-guided delimited matrix: {error}"),
}
})?;
let stream_ms = t_stream.elapsed().as_secs_f64() * 1000.0;
if stream_ms > 100.0 {
log::info!(
"[DATA-LOAD] delim_schema_direct | n_rows={} | n_cols={} | {:.1}ms",
total_rows,
p,
stream_ms
);
}
let column_kinds = col_meta.iter().map(|meta| meta.kind).collect();
let schema_out = DataSchema {
columns: col_meta.into_iter().map(|meta| meta.schema_col).collect(),
};
return Ok(EncodedDataset {
headers,
values,
schema: schema_out,
column_kinds,
});
}
let mut inference = vec![DelimitedInferenceState::default(); p];
let mut total_rows: usize = 0;
let t_stream = std::time::Instant::now();
let mut record = StringRecord::new();
while rdr
.read_record(&mut record)
.map_err(|e| DataError::ParseError {
reason: format!("failed reading row: {e}"),
})?
{
if record.len() != all_headers.len() {
return Err(DataError::SchemaMismatch {
reason: format!(
"row width mismatch at row {}: got {} fields, expected {}",
total_rows + 1,
record.len(),
all_headers.len()
),
});
}
total_rows += 1;
for j in 0..p {
let raw = record
.get(selected_indices[j])
.expect("record width was checked against the header row above")
.trim();
if needs_inference[j] {
inference[j].observe(raw, total_rows, &headers[j])?;
} else {
parse_cell_with_schema(raw, &col_meta[j], total_rows, &headers[j], &unseen_policy)?;
}
}
}
let stream_ms = t_stream.elapsed().as_secs_f64() * 1000.0;
if stream_ms > 100.0 {
let n_inf = needs_inference.iter().filter(|x| **x).count();
log::info!(
"[DATA-LOAD] delim_schema_stream | n_rows={} | n_cols={} | n_inf={} | {:.1}ms",
total_rows,
p,
n_inf,
stream_ms
);
}
if total_rows == 0 {
return Err(DataError::EmptyInput {
reason: "file has no rows".to_string(),
});
}
let t_finalize = std::time::Instant::now();
for j in 0..p {
if needs_inference[j] {
let kind = inference[j].kind(false);
col_meta[j].kind = kind;
col_meta[j].schema_col.kind = kind;
}
}
let finalize_ms = t_finalize.elapsed().as_secs_f64() * 1000.0;
if finalize_ms > 100.0 {
log::info!(
"[DATA-LOAD] delim_schema_finalize | n_cols={} | {:.1}ms",
p,
finalize_ms
);
}
let t_assemble = std::time::Instant::now();
let mut values = Array2::<f64>::zeros((total_rows, p));
let mut inferred_encoders = (0..p)
.map(|j| {
(needs_inference[j] && matches!(col_meta[j].kind, ColumnKindTag::Categorical))
.then(CategoricalEncoder::default)
})
.collect::<Vec<_>>();
let mut encode_rdr = ReaderBuilder::new()
.has_headers(true)
.delimiter(delimiter)
.from_path(path)
.map_err(|e| DataError::ParseError {
reason: format!("failed to reopen '{}': {e}", path.display()),
})?;
encode_rdr.headers().map_err(|e| DataError::ParseError {
reason: format!("failed to reread headers: {e}"),
})?;
let mut encoded_rows = 0usize;
while encode_rdr
.read_record(&mut record)
.map_err(|e| DataError::ParseError {
reason: format!("failed reading row: {e}"),
})?
{
if record.len() != all_headers.len() {
return Err(DataError::SchemaMismatch {
reason: format!(
"row width mismatch at row {}: got {} fields, expected {}",
encoded_rows + 1,
record.len(),
all_headers.len()
),
});
}
if encoded_rows >= total_rows {
return Err(DataError::SchemaMismatch {
reason: "data file changed while its schema was being discovered".to_string(),
});
}
for j in 0..p {
let raw = record
.get(selected_indices[j])
.expect("record width was checked against the header row above")
.trim();
values[[encoded_rows, j]] = if !needs_inference[j] {
parse_cell_with_schema(
raw,
&col_meta[j],
encoded_rows + 1,
&headers[j],
&unseen_policy,
)?
} else {
match col_meta[j].kind {
ColumnKindTag::Continuous | ColumnKindTag::Binary => {
parse_inferred_numeric_cell(raw, encoded_rows + 1, &headers[j])?
}
ColumnKindTag::Categorical => {
if raw.is_empty() {
return Err(DataError::EmptyInput {
reason: format!(
"empty field at row {}, column '{}'",
encoded_rows + 1,
&headers[j]
),
});
}
let encoder = inferred_encoders[j]
.as_mut()
.expect("inferred categorical encoder");
encoder.encode(raw) as f64
}
}
};
}
encoded_rows += 1;
}
if encoded_rows != total_rows {
return Err(DataError::SchemaMismatch {
reason: "data file changed while its schema was being discovered".to_string(),
});
}
for (j, encoder) in inferred_encoders.into_iter().enumerate() {
if let Some(encoder) = encoder {
col_meta[j].schema_col.levels =
encoder.finish(values.column_mut(j), LevelOrder::Canonical);
}
}
let assemble_ms = t_assemble.elapsed().as_secs_f64() * 1000.0;
if assemble_ms > 100.0 {
log::info!(
"[DATA-LOAD] delim_schema_assemble | n_rows={} | n_cols={} | {:.1}ms",
total_rows,
p,
assemble_ms
);
}
let column_kinds = col_meta.iter().map(|meta| meta.kind).collect();
let schema_out = DataSchema {
columns: col_meta.into_iter().map(|m| m.schema_col).collect(),
};
Ok(EncodedDataset {
headers,
values,
schema: schema_out,
column_kinds,
})
}
fn parse_cell_with_schema(
raw: &str,
meta: &ColMeta<'_>,
row: usize,
col_name: &str,
unseen_policy: &UnseenCategoryPolicy,
) -> Result<f64, DataError> {
let val = match meta.kind {
ColumnKindTag::Continuous if is_missing_marker(raw) => f64::NAN,
ColumnKindTag::Continuous => raw.parse::<f64>().map_err(|err| {
DataError::SchemaMismatch {
reason: format!(
"column '{}' is continuous in schema but row {} has non-numeric value '{}': {}",
col_name, row, raw, err
),
}
})?,
ColumnKindTag::Binary if is_missing_marker(raw) => f64::NAN,
ColumnKindTag::Binary => {
let v = raw
.parse::<f64>()
.map_err(|err| DataError::SchemaMismatch {
reason: format!(
"column '{}' is binary in schema but row {} has non-numeric value '{}': {}",
col_name, row, raw, err
),
})?;
if (v - 0.0).abs() >= 1e-12 && (v - 1.0).abs() >= 1e-12 {
return Err(DataError::SchemaMismatch {
reason: format!(
"column '{}' is binary in schema but row {} has value {}; expected 0 or 1",
col_name, row, v
),
});
}
v
}
ColumnKindTag::Categorical => {
let map = meta
.level_map
.as_ref()
.ok_or_else(|| DataError::EncodingFailure {
reason: "internal categorical schema map missing".to_string(),
})?;
match map.get(raw) {
Some(v) => *v,
None => unseen_policy
.unseen_code_for(col_name, meta.schema_col.levels.len())
.ok_or_else(|| DataError::SchemaMismatch {
reason: format!(
"unseen level '{}' in categorical column '{}' at row {}",
raw, col_name, row
),
})?,
}
}
};
if !val.is_finite() && !is_missing_marker(raw) {
return Err(DataError::InvalidValue {
reason: format!("non-finite value at row {}, column '{}'", row, col_name),
});
}
Ok(val)
}
struct ColMeta<'a> {
kind: ColumnKindTag,
level_map: Option<HashMap<&'a str, f64>>,
schema_col: SchemaColumn,
}
fn arrow_field_is_string(dt: &arrow::datatypes::DataType) -> bool {
use arrow::datatypes::DataType;
match dt {
DataType::Utf8 | DataType::LargeUtf8 => true,
DataType::Dictionary(_, value_type) => arrow_field_is_string(value_type),
_ => false,
}
}
fn write_arrow_numeric_values(
values: impl IntoIterator<Item = Option<f64>>,
mut output: ArrayViewMut1<'_, f64>,
categorical_encoder: Option<&mut CategoricalEncoder>,
all_binary: &mut bool,
saw_numeric: &mut bool,
) {
match categorical_encoder {
Some(encoder) => {
for (batch_row, value) in values.into_iter().enumerate() {
output[batch_row] = match value.filter(|value| value.is_finite()) {
Some(value) => encoder.encode(&value.to_string()) as f64,
None => f64::NAN,
};
}
}
None => {
for (batch_row, value) in values.into_iter().enumerate() {
let Some(value) = value.filter(|value| value.is_finite()) else {
output[batch_row] = f64::NAN;
continue;
};
*saw_numeric = true;
if (value - 0.0).abs() >= 1e-12 && (value - 1.0).abs() >= 1e-12 {
*all_binary = false;
}
output[batch_row] = value;
}
}
}
}
fn arrow_dictionary_string_value_at<'a, K>(
col: &'a dyn arrow::array::Array,
index: usize,
logical_row: usize,
header: &str,
) -> Result<Option<&'a str>, DataError>
where
K: arrow::datatypes::ArrowDictionaryKeyType,
{
use arrow::array::DictionaryArray;
let dictionary = col
.as_any()
.downcast_ref::<DictionaryArray<K>>()
.ok_or_else(|| DataError::EncodingFailure {
reason: format!(
"Arrow dictionary column '{}' did not match its declared key type",
header
),
})?;
let Some(value_index) = dictionary.key(index) else {
return Ok(None);
};
if value_index >= dictionary.values().len() {
return Err(DataError::EncodingFailure {
reason: format!(
"Arrow dictionary column '{}' has out-of-range key {} at row {}",
header, value_index, logical_row
),
});
}
arrow_string_value_at(
dictionary.values().as_ref(),
value_index,
logical_row,
header,
)
}
fn arrow_string_value_at<'a>(
col: &'a dyn arrow::array::Array,
index: usize,
logical_row: usize,
header: &str,
) -> Result<Option<&'a str>, DataError> {
use arrow::array::{LargeStringArray, StringArray};
use arrow::datatypes::{
DataType, Int8Type, Int16Type, Int32Type, Int64Type, UInt8Type, UInt16Type, UInt32Type,
UInt64Type,
};
if index >= col.len() {
return Err(DataError::EncodingFailure {
reason: format!(
"Arrow string column '{}' has out-of-range index {} at row {}",
header, index, logical_row
),
});
}
if col.is_null(index) {
return Ok(None);
}
match col.data_type() {
DataType::Utf8 => col
.as_any()
.downcast_ref::<StringArray>()
.map(|array| Some(array.value(index)))
.ok_or_else(|| DataError::EncodingFailure {
reason: format!("Arrow column '{}' could not be read as Utf8", header),
}),
DataType::LargeUtf8 => col
.as_any()
.downcast_ref::<LargeStringArray>()
.map(|array| Some(array.value(index)))
.ok_or_else(|| DataError::EncodingFailure {
reason: format!("Arrow column '{}' could not be read as LargeUtf8", header),
}),
DataType::Dictionary(key_type, _) => match key_type.as_ref() {
DataType::Int8 => {
arrow_dictionary_string_value_at::<Int8Type>(col, index, logical_row, header)
}
DataType::Int16 => {
arrow_dictionary_string_value_at::<Int16Type>(col, index, logical_row, header)
}
DataType::Int32 => {
arrow_dictionary_string_value_at::<Int32Type>(col, index, logical_row, header)
}
DataType::Int64 => {
arrow_dictionary_string_value_at::<Int64Type>(col, index, logical_row, header)
}
DataType::UInt8 => {
arrow_dictionary_string_value_at::<UInt8Type>(col, index, logical_row, header)
}
DataType::UInt16 => {
arrow_dictionary_string_value_at::<UInt16Type>(col, index, logical_row, header)
}
DataType::UInt32 => {
arrow_dictionary_string_value_at::<UInt32Type>(col, index, logical_row, header)
}
DataType::UInt64 => {
arrow_dictionary_string_value_at::<UInt64Type>(col, index, logical_row, header)
}
other => Err(DataError::InvalidValue {
reason: format!(
"unsupported Arrow dictionary key type {:?} for column '{}'",
other, header
),
}),
},
other => Err(DataError::InvalidValue {
reason: format!(
"unsupported Arrow string column type {:?} for column '{}'",
other, header
),
}),
}
}
fn decode_arrow_batch_column_into(
col: &dyn arrow::array::Array,
base_row: usize,
header: &str,
is_string_col: bool,
mut output: ArrayViewMut1<'_, f64>,
mut categorical_encoder: Option<&mut CategoricalEncoder>,
all_binary: &mut bool,
saw_numeric: &mut bool,
) -> Result<(), DataError> {
use arrow::array::{
Array as _, BooleanArray, Float32Array, Float64Array, Int8Array, Int16Array, Int32Array,
Int64Array, UInt8Array, UInt16Array, UInt32Array, UInt64Array,
};
use arrow::datatypes::DataType;
let n_rows = output.len();
if col.len() != n_rows {
return Err(DataError::SchemaMismatch {
reason: format!(
"Arrow column '{}' has {} rows, but its record batch has {}",
header,
col.len(),
n_rows
),
});
}
if is_string_col {
let encoder =
categorical_encoder
.as_deref_mut()
.ok_or_else(|| DataError::EncodingFailure {
reason: format!("categorical Arrow encoder missing for column '{header}'"),
})?;
for batch_row in 0..n_rows {
output[batch_row] = match arrow_string_value_at(
col,
batch_row,
base_row + batch_row + 1,
header,
)? {
Some("") | None => f64::NAN,
Some(label) => encoder.encode(label) as f64,
};
}
return Ok(());
}
let decoded_col;
let col: &dyn arrow::array::Array = if let DataType::Dictionary(_, value_type) = col.data_type()
{
decoded_col = arrow::compute::cast(col, value_type).map_err(|e| DataError::ParseError {
reason: format!(
"failed to decode dictionary-encoded numeric column '{}': {e}",
header
),
})?;
decoded_col.as_ref()
} else {
col
};
macro_rules! write_primitive {
($array_type:ty, $convert:expr) => {{
let array = col
.as_any()
.downcast_ref::<$array_type>()
.expect("array type is the one this `col.data_type()` arm matched");
write_arrow_numeric_values(
(0..n_rows).map(|index| {
(!array.is_null(index)).then(|| $convert(array.value(index)))
}),
output,
categorical_encoder.as_deref_mut(),
all_binary,
saw_numeric,
);
Ok(())
}};
}
match col.data_type() {
DataType::Float64 => write_primitive!(Float64Array, |value: f64| value),
DataType::Float32 => write_primitive!(Float32Array, |value: f32| value as f64),
DataType::Int64 => write_primitive!(Int64Array, |value: i64| value as f64),
DataType::Int32 => write_primitive!(Int32Array, |value: i32| value as f64),
DataType::Int16 => write_primitive!(Int16Array, |value: i16| value as f64),
DataType::Int8 => write_primitive!(Int8Array, |value: i8| value as f64),
DataType::UInt64 => write_primitive!(UInt64Array, |value: u64| value as f64),
DataType::UInt32 => write_primitive!(UInt32Array, |value: u32| value as f64),
DataType::UInt16 => write_primitive!(UInt16Array, |value: u16| value as f64),
DataType::UInt8 => write_primitive!(UInt8Array, |value: u8| value as f64),
DataType::Boolean => {
let arr = col
.as_any()
.downcast_ref::<BooleanArray>()
.expect("array type is BooleanArray in the DataType::Boolean arm");
write_arrow_numeric_values(
(0..n_rows).map(|index| {
(!arr.is_null(index)).then(|| if arr.value(index) { 1.0 } else { 0.0 })
}),
output,
categorical_encoder.as_deref_mut(),
all_binary,
saw_numeric,
);
Ok(())
}
other => Err(DataError::InvalidValue {
reason: format!(
"unsupported Arrow column type {:?} for column '{}'",
other, header
),
}),
}
}
pub fn encode_arrow_record_batch_reader_with_inferred_schema(
reader: &mut dyn arrow::record_batch::RecordBatchReader,
headers: Vec<String>,
) -> Result<EncodedDataset, DataError> {
if headers.is_empty() {
return Err(DataError::EmptyInput {
reason: "Arrow table must have at least one header column".to_string(),
});
}
let mut seen_headers = HashSet::<&str>::with_capacity(headers.len());
for (column, header) in headers.iter().enumerate() {
if header.trim().is_empty() {
return Err(DataError::EmptyInput {
reason: format!("Arrow header at column {} cannot be empty", column + 1),
});
}
if !seen_headers.insert(header.as_str()) {
return Err(DataError::SchemaMismatch {
reason: format!("duplicate Arrow header '{}'", header),
});
}
}
let arrow_schema = reader.schema();
let p = headers.len();
if arrow_schema.fields().len() != p {
return Err(DataError::SchemaMismatch {
reason: format!(
"Arrow schema has {} columns, but {} normalized headers were supplied",
arrow_schema.fields().len(),
p
),
});
}
let is_string_col = arrow_schema
.fields()
.iter()
.map(|field| arrow_field_is_string(field.data_type()))
.collect::<Vec<_>>();
let mut all_binary = vec![true; p];
let mut saw_numeric = vec![false; p];
let mut categorical_encoders = is_string_col
.iter()
.map(|&is_string| is_string.then(CategoricalEncoder::default))
.collect::<Vec<_>>();
let mut encoded_values = Vec::<f64>::new();
let mut rows_seen = 0usize;
for batch_result in reader {
let batch = batch_result.map_err(|error| DataError::ParseError {
reason: format!("failed to read Arrow record batch: {error}"),
})?;
if batch.num_columns() != p {
return Err(DataError::SchemaMismatch {
reason: format!(
"Arrow record batch has {} columns, but {} normalized headers were supplied",
batch.num_columns(),
p
),
});
}
for j in 0..p {
let expected = arrow_schema.field(j).data_type();
let actual = batch.column(j).data_type();
if actual != expected {
return Err(DataError::SchemaMismatch {
reason: format!(
"Arrow column '{}' changed type between schema and batch: expected {:?}, got {:?}",
headers[j], expected, actual
),
});
}
}
let n_rows = batch.num_rows();
let batch_values = n_rows
.checked_mul(p)
.ok_or_else(|| DataError::EncodingFailure {
reason: "Arrow batch dimensions do not fit in memory address space".to_string(),
})?;
let next_len = encoded_values
.len()
.checked_add(batch_values)
.ok_or_else(|| DataError::EncodingFailure {
reason: "Arrow dataset dimensions do not fit in memory address space".to_string(),
})?;
encoded_values
.try_reserve(batch_values)
.map_err(|error| DataError::EncodingFailure {
reason: format!("failed to reserve Arrow dataset storage: {error}"),
})?;
let batch_offset = encoded_values.len();
encoded_values.resize(next_len, 0.0);
let mut batch_output = ndarray::ArrayViewMut2::from_shape(
(n_rows, p),
&mut encoded_values[batch_offset..next_len],
)
.map_err(|error| DataError::EncodingFailure {
reason: format!("failed to shape Arrow batch output: {error}"),
})?;
let decoded_columns = batch_output
.axis_iter_mut(Axis(1))
.into_par_iter()
.zip(categorical_encoders.par_iter_mut())
.zip(all_binary.par_iter_mut())
.zip(saw_numeric.par_iter_mut())
.enumerate()
.map(|(j, (((output, encoder), column_all_binary), column_saw_numeric))| {
decode_arrow_batch_column_into(
batch.column(j).as_ref(),
rows_seen,
&headers[j],
is_string_col[j],
output,
encoder.as_mut(),
column_all_binary,
column_saw_numeric,
)
})
.collect::<Vec<_>>();
for decoded in decoded_columns {
decoded?;
}
rows_seen = rows_seen
.checked_add(n_rows)
.ok_or_else(|| DataError::EncodingFailure {
reason: "Arrow row count does not fit in memory address space".to_string(),
})?;
}
if rows_seen == 0 {
return Err(DataError::EmptyInput {
reason: "Arrow table data cannot be empty".to_string(),
});
}
let mut values = Array2::from_shape_vec((rows_seen, p), encoded_values).map_err(|error| {
DataError::EncodingFailure {
reason: format!("failed to shape encoded Arrow dataset: {error}"),
}
})?;
let mut levels = vec![Vec::<String>::new(); p];
for (j, encoder) in categorical_encoders.into_iter().enumerate() {
if let Some(encoder) = encoder {
levels[j] = encoder.finish(values.column_mut(j), LevelOrder::Canonical);
}
}
let mut schema_columns = Vec::<SchemaColumn>::with_capacity(p);
let mut column_kinds = Vec::<ColumnKindTag>::with_capacity(p);
for (j, name) in headers.iter().enumerate() {
let kind = if is_string_col[j] {
ColumnKindTag::Categorical
} else if all_binary[j] && saw_numeric[j] {
ColumnKindTag::Binary
} else {
ColumnKindTag::Continuous
};
column_kinds.push(kind);
schema_columns.push(SchemaColumn {
name: name.clone(),
kind,
levels: std::mem::take(&mut levels[j]),
});
}
Ok(EncodedDataset {
headers,
values,
schema: DataSchema {
columns: schema_columns,
},
column_kinds,
})
}
fn load_parquet_inferred(
path: &Path,
requested_columns: &[String],
categorical_roles: &HashSet<&str>,
) -> Result<EncodedDataset, DataError> {
use parquet::arrow::{ProjectionMask, arrow_reader::ParquetRecordBatchReaderBuilder};
use rayon::prelude::*;
use std::fs::File;
let t_open = std::time::Instant::now();
let file = File::open(path).map_err(|e| DataError::ParseError {
reason: format!("failed to open parquet '{}': {e}", path.display()),
})?;
let builder =
ParquetRecordBatchReaderBuilder::try_new(file).map_err(|e| DataError::ParseError {
reason: format!("failed to read parquet metadata '{}': {e}", path.display()),
})?;
let full_schema = builder.schema().clone();
let all_headers: Vec<String> = full_schema
.fields()
.iter()
.map(|f| f.name().clone())
.collect();
if all_headers.is_empty() {
return Err(DataError::EmptyInput {
reason: "parquet file has no columns".to_string(),
});
}
let selected_indices = resolve_requested_columns(&all_headers, requested_columns)?;
let headers = projected_headers(&all_headers, &selected_indices);
let selected_fields = selected_indices
.iter()
.map(|&idx| full_schema.fields()[idx].clone())
.collect::<Vec<_>>();
let total_rows =
usize::try_from(builder.metadata().file_metadata().num_rows()).map_err(|_| {
DataError::ParseError {
reason: "parquet row count does not fit in memory address space".to_string(),
}
})?;
if total_rows == 0 {
return Err(DataError::EmptyInput {
reason: "parquet file has no rows".to_string(),
});
}
let projection =
ProjectionMask::roots(builder.parquet_schema(), selected_indices.iter().copied());
let reader =
builder
.with_projection(projection)
.build()
.map_err(|e| DataError::ParseError {
reason: format!("failed to build parquet reader: {e}"),
})?;
let p = headers.len();
let open_ms = t_open.elapsed().as_secs_f64() * 1000.0;
if open_ms > 100.0 {
log::info!(
"[DATA-LOAD] parquet_open+meta | n_headers={} | n_proj={} | {:.1}ms",
all_headers.len(),
p,
open_ms
);
}
let t_batches = std::time::Instant::now();
let is_string_col = selected_fields
.iter()
.map(|field| arrow_field_is_string(field.data_type()))
.collect::<Vec<_>>();
let forced_numeric_categorical = headers
.iter()
.enumerate()
.map(|(j, header)| !is_string_col[j] && categorical_roles.contains(header.as_str()))
.collect::<Vec<_>>();
let mut values = Array2::<f64>::zeros((total_rows, p));
let mut all_binary = vec![true; p];
let mut saw_numeric = vec![false; p];
let mut categorical_encoders = (0..p)
.map(|j| {
(is_string_col[j] || forced_numeric_categorical[j]).then(CategoricalEncoder::default)
})
.collect::<Vec<_>>();
let mut rows_seen = 0usize;
for batch_result in reader {
let batch = batch_result.map_err(|e| DataError::ParseError {
reason: format!("failed to read parquet record batch: {e}"),
})?;
let n_rows = batch.num_rows();
if rows_seen.saturating_add(n_rows) > total_rows {
return Err(DataError::SchemaMismatch {
reason: "parquet row count changed while reading record batches".to_string(),
});
}
let decoded_columns = values
.slice_mut(s![rows_seen..rows_seen + n_rows, ..])
.axis_iter_mut(Axis(1))
.into_par_iter()
.zip(categorical_encoders.par_iter_mut())
.zip(all_binary.par_iter_mut())
.zip(saw_numeric.par_iter_mut())
.enumerate()
.map(|(j, (((output, encoder), column_all_binary), column_saw_numeric))| {
decode_arrow_batch_column_into(
batch.column(j).as_ref(),
rows_seen,
&headers[j],
is_string_col[j],
output,
encoder.as_mut(),
column_all_binary,
column_saw_numeric,
)
})
.collect::<Vec<_>>();
for decoded in decoded_columns {
decoded?;
}
rows_seen += n_rows;
}
if rows_seen != total_rows {
return Err(DataError::SchemaMismatch {
reason: format!(
"parquet metadata reports {total_rows} rows but record batches yielded {rows_seen}"
),
});
}
let batches_ms = t_batches.elapsed().as_secs_f64() * 1000.0;
if batches_ms > 100.0 {
log::info!(
"[DATA-LOAD] parquet_batches_decode | n_rows={} | n_cols={} | {:.1}ms",
total_rows,
p,
batches_ms
);
}
let t_schema = std::time::Instant::now();
let mut levels = vec![Vec::<String>::new(); p];
for (j, encoder) in categorical_encoders.into_iter().enumerate() {
if let Some(encoder) = encoder {
let order = if forced_numeric_categorical[j] {
LevelOrder::Canonical
} else {
LevelOrder::Encounter
};
levels[j] = encoder.finish(values.column_mut(j), order);
}
}
let mut schema_cols = Vec::<SchemaColumn>::with_capacity(p);
let mut column_kinds = Vec::<ColumnKindTag>::with_capacity(p);
for j in 0..p {
let kind = if is_string_col[j] || forced_numeric_categorical[j] {
ColumnKindTag::Categorical
} else if all_binary[j] && saw_numeric[j] {
ColumnKindTag::Binary
} else {
ColumnKindTag::Continuous
};
column_kinds.push(kind);
schema_cols.push(SchemaColumn {
name: headers[j].clone(),
kind,
levels: std::mem::take(&mut levels[j]),
});
}
let schema_ms = t_schema.elapsed().as_secs_f64() * 1000.0;
if schema_ms > 100.0 {
let n_cat = column_kinds
.iter()
.filter(|k| matches!(k, ColumnKindTag::Categorical))
.count();
log::info!(
"[DATA-LOAD] parquet_finalize_schema | n_cols={} | n_cat={} | {:.1}ms",
p,
n_cat,
schema_ms
);
}
Ok(EncodedDataset {
headers,
values,
schema: DataSchema {
columns: schema_cols,
},
column_kinds,
})
}
fn load_parquet_with_schema(
path: &Path,
schema: &DataSchema,
unseen_policy: UnseenCategoryPolicy,
requested_columns: &[String],
) -> Result<EncodedDataset, DataError> {
let inferred = load_parquet_inferred(path, requested_columns, &HashSet::new())?;
let p = inferred.headers.len();
let n = inferred.values.nrows();
let schema_byname: HashMap<&str, &SchemaColumn> = schema
.columns
.iter()
.map(|c| (c.name.as_str(), c))
.collect();
let mut column_kinds = Vec::<ColumnKindTag>::with_capacity(p);
let mut schema_cols = Vec::<SchemaColumn>::with_capacity(p);
let mut values = inferred.values;
for j in 0..p {
let name = &inferred.headers[j];
if let Some(sc) = schema_byname.get(name.as_str()) {
column_kinds.push(sc.kind);
schema_cols.push((*sc).clone());
match sc.kind {
ColumnKindTag::Continuous => {
if matches!(inferred.column_kinds[j], ColumnKindTag::Categorical) {
return Err(DataError::SchemaMismatch {
reason: format!(
"column '{}' is continuous in schema but parquet column is string/categorical",
name
),
});
}
}
ColumnKindTag::Binary => {
if matches!(inferred.column_kinds[j], ColumnKindTag::Categorical) {
return Err(DataError::SchemaMismatch {
reason: format!(
"column '{}' is binary in schema but parquet column is string/categorical",
name
),
});
}
if let Some(row) = values.column(j).iter().position(|value| {
value.is_finite()
&& (*value - 0.0).abs() >= 1e-12
&& (*value - 1.0).abs() >= 1e-12
}) {
return Err(DataError::SchemaMismatch {
reason: format!(
"column '{}' is binary in schema but row {} has value {}; expected 0 or 1",
name,
row + 1,
values[[row, j]]
),
});
}
}
ColumnKindTag::Categorical => {
if !matches!(inferred.column_kinds[j], ColumnKindTag::Categorical) {
return Err(DataError::SchemaMismatch {
reason: format!(
"column '{}' is categorical in schema but parquet column is numeric",
name
),
});
}
let inferred_col = &inferred.schema.columns[j];
let schema_level_map: HashMap<&str, f64> = sc
.levels
.iter()
.enumerate()
.map(|(idx, v)| (v.as_str(), idx as f64))
.collect();
let inferred_to_schema: Vec<f64> = inferred_col
.levels
.iter()
.map(|lv| {
schema_level_map
.get(lv.as_str())
.copied()
.or_else(|| unseen_policy.unseen_code_for(name, sc.levels.len()))
.ok_or_else(|| DataError::SchemaMismatch {
reason: format!(
"unseen level '{}' in categorical column '{}'",
lv, name
),
})
})
.collect::<Result<Vec<_>, _>>()?;
for i in 0..n {
let old_code = values[[i, j]] as usize;
if old_code >= inferred_to_schema.len() {
let Some(unseen_code) =
unseen_policy.unseen_code_for(name, sc.levels.len())
else {
return Err(DataError::SchemaMismatch {
reason: format!(
"unseen categorical code at row {}, column '{}'",
i + 1,
name
),
});
};
values[[i, j]] = unseen_code;
continue;
}
values[[i, j]] = inferred_to_schema[old_code];
}
}
}
} else {
column_kinds.push(inferred.column_kinds[j]);
schema_cols.push(inferred.schema.columns[j].clone());
}
}
Ok(EncodedDataset {
headers: inferred.headers,
values,
schema: DataSchema {
columns: schema_cols,
},
column_kinds,
})
}
pub fn encode_recordswith_inferred_schema(
headers: Vec<String>,
records: Vec<StringRecord>,
) -> Result<EncodedDataset, String> {
if records.is_empty() {
return Err(DataError::EmptyInput {
reason: "table data cannot be empty".to_string(),
}
.into());
}
let schema_cols = headers
.par_iter()
.enumerate()
.map(|(j, name)| infer_schema_column(name, &records, j).map_err(String::from))
.collect::<Result<Vec<SchemaColumn>, String>>()?;
let schema = DataSchema {
columns: schema_cols,
};
encode_recordswith_schema(headers, records, &schema, UnseenCategoryPolicy::Error)
}
pub fn encode_recordswith_schema(
headers: Vec<String>,
records: Vec<StringRecord>,
schema: &DataSchema,
unseen_policy: UnseenCategoryPolicy,
) -> Result<EncodedDataset, String> {
let n = records.len();
if n == 0 {
return Err(DataError::EmptyInput {
reason: "table data cannot be empty".to_string(),
}
.into());
}
let p = headers.len();
if p == 0 {
return Err(DataError::EmptyInput {
reason: "table data must have at least one header column".to_string(),
}
.into());
}
for (i, rec) in records.iter().enumerate() {
if rec.len() != p {
return Err(DataError::SchemaMismatch {
reason: format!(
"row width mismatch at row {}: got {} fields, expected {} (one per header)",
i + 1,
rec.len(),
p
),
}
.into());
}
}
let schema_byname: HashMap<&str, &SchemaColumn> = schema
.columns
.iter()
.map(|c| (c.name.as_str(), c))
.collect();
let encoded_columns = headers
.par_iter()
.enumerate()
.map(|(j, name)| {
let inferred_for_extra;
let col_schema = if let Some(s) = schema_byname.get(name.as_str()) {
*s
} else {
inferred_for_extra =
infer_schema_column(name, &records, j).map_err(String::from)?;
&inferred_for_extra
};
let column = encode_one_column(name, &records, j, col_schema, &unseen_policy)?;
Ok::<(ColumnKindTag, Vec<f64>), String>((col_schema.kind, column))
})
.collect::<Result<Vec<(ColumnKindTag, Vec<f64>)>, String>>()?;
let mut column_kinds = Vec::<ColumnKindTag>::with_capacity(p);
let mut values = Array2::<f64>::zeros((n, p));
for (j, (kind, column)) in encoded_columns.into_iter().enumerate() {
column_kinds.push(kind);
values
.column_mut(j)
.assign(&ndarray::ArrayView1::from(&column));
}
Ok(EncodedDataset {
headers,
values,
schema: schema.clone(),
column_kinds,
})
}
fn encode_one_column(
name: &str,
records: &[StringRecord],
j: usize,
col_schema: &SchemaColumn,
unseen_policy: &UnseenCategoryPolicy,
) -> Result<Vec<f64>, String> {
let level_map = if matches!(col_schema.kind, ColumnKindTag::Categorical) {
Some(
col_schema
.levels
.iter()
.enumerate()
.map(|(idx, v)| (v.as_str(), idx as f64))
.collect::<HashMap<_, _>>(),
)
} else {
None
};
let mut column = Vec::<f64>::with_capacity(records.len());
for (i, rec) in records.iter().enumerate() {
let raw = rec
.get(j)
.ok_or_else(|| {
String::from(DataError::SchemaMismatch {
reason: format!("missing field at row {}, col {}", i + 1, j + 1),
})
})?
.trim();
if raw.is_empty() {
return Err(DataError::EmptyInput {
reason: format!("empty field at row {}, column '{}'", i + 1, name),
}
.into());
}
let val = match col_schema.kind {
ColumnKindTag::Continuous if is_missing_marker(raw) => f64::NAN,
ColumnKindTag::Continuous => raw.parse::<f64>().map_err(|err| {
String::from(DataError::SchemaMismatch {
reason: format!(
"column '{}' is continuous in schema but row {} has non-numeric value '{}': {}",
name,
i + 1,
raw,
err
),
})
})?,
ColumnKindTag::Binary if is_missing_marker(raw) => f64::NAN,
ColumnKindTag::Binary => {
let v = raw.parse::<f64>().map_err(|err| {
String::from(DataError::SchemaMismatch {
reason: format!(
"column '{}' is binary in schema but row {} has non-numeric value '{}': {}",
name,
i + 1,
raw,
err
),
})
})?;
if (v - 0.0).abs() >= 1e-12 && (v - 1.0).abs() >= 1e-12 {
return Err(DataError::SchemaMismatch {
reason: format!(
"column '{}' is binary in schema but row {} has value {}; expected 0 or 1",
name,
i + 1,
v
),
}
.into());
}
v
}
ColumnKindTag::Categorical => {
let map = level_map.as_ref().ok_or_else(|| {
String::from(DataError::EncodingFailure {
reason: "internal categorical schema map missing".to_string(),
})
})?;
match map.get(raw) {
Some(v) => *v,
None => unseen_policy
.unseen_code_for(name, col_schema.levels.len())
.ok_or_else(|| {
String::from(DataError::SchemaMismatch {
reason: format!(
"unseen level '{}' in categorical column '{}' at row {}; allowed levels: {}",
raw,
name,
i + 1,
col_schema.levels.join(",")
),
})
})?,
}
}
};
if !val.is_finite() && !is_missing_marker(raw) {
return Err(DataError::InvalidValue {
reason: format!("non-finite value at row {}, column '{}'", i + 1, name),
}
.into());
}
column.push(val);
}
Ok(column)
}
fn infer_schema_column(
name: &str,
records: &[StringRecord],
col_idx: usize,
) -> Result<SchemaColumn, DataError> {
let mut all_numeric = true;
let mut all_binary = true;
let mut saw_numeric = false;
let mut levels = Vec::<String>::new();
let mut level_index = HashMap::<String, usize>::new();
let mut missing_markers = Vec::<String>::new();
for (i, rec) in records.iter().enumerate() {
let raw = rec
.get(col_idx)
.ok_or_else(|| DataError::SchemaMismatch {
reason: format!("missing field at row {}, col {}", i + 1, col_idx + 1),
})?
.trim();
if raw.is_empty() {
return Err(DataError::EmptyInput {
reason: format!("empty field at row {}, column '{}'", i + 1, name),
});
}
if is_missing_marker(raw) {
missing_markers.push(raw.to_string());
continue;
}
if let Ok(v) = raw.parse::<f64>() {
saw_numeric = true;
if !v.is_finite() {
return Err(DataError::InvalidValue {
reason: format!("non-finite value at row {}, column '{}'", i + 1, name),
});
}
if (v - 0.0).abs() >= 1e-12 && (v - 1.0).abs() >= 1e-12 {
all_binary = false;
}
} else {
all_numeric = false;
all_binary = false;
level_index.entry(raw.to_string()).or_insert_with(|| {
let idx = levels.len();
levels.push(raw.to_string());
idx
});
}
}
let numeric_column = all_numeric && saw_numeric;
if !numeric_column {
for marker in missing_markers {
level_index.entry(marker.clone()).or_insert_with(|| {
let idx = levels.len();
levels.push(marker);
idx
});
}
}
let kind = if numeric_column {
if all_binary {
ColumnKindTag::Binary
} else {
ColumnKindTag::Continuous
}
} else {
ColumnKindTag::Categorical
};
if matches!(kind, ColumnKindTag::Categorical) {
sort_levels_canonical(&mut levels);
}
Ok(SchemaColumn {
name: name.to_string(),
kind,
levels: if matches!(kind, ColumnKindTag::Categorical) {
levels
} else {
Vec::new()
},
})
}
pub fn infer_and_encode_column_major(
name: &str,
column: &[&str],
col_index: usize,
) -> Result<(SchemaColumn, Vec<f64>), String> {
if column.is_empty() {
return Err(DataError::EmptyInput {
reason: "table data cannot be empty".to_string(),
}
.into());
}
let force_categorical = column.iter().any(|c| strip_categorical_sentinel(c).1);
let mut all_numeric = !force_categorical;
let mut all_binary = !force_categorical;
let mut levels = Vec::<String>::new();
let mut level_index = HashMap::<String, usize>::new();
let mut trimmed = Vec::<&str>::with_capacity(column.len());
let mut parsed = Vec::<Option<f64>>::with_capacity(column.len());
let mut saw_numeric = false;
let mut missing_positions = Vec::<usize>::new();
for (i, raw_field) in column.iter().enumerate() {
let (raw, _) = strip_categorical_sentinel(raw_field);
let raw = raw.trim();
if raw.is_empty() {
return Err(DataError::EmptyInput {
reason: format!("empty field at row {}, column '{}'", i + 1, name),
}
.into());
}
if !force_categorical {
if is_missing_marker(raw) {
missing_positions.push(i);
parsed.push(Some(f64::NAN));
trimmed.push(raw);
continue;
}
if let Ok(v) = raw.parse::<f64>() {
saw_numeric = true;
if !v.is_finite() {
return Err(DataError::InvalidValue {
reason: format!("non-finite value at row {}, column '{}'", i + 1, name),
}
.into());
}
if (v - 0.0).abs() >= 1e-12 && (v - 1.0).abs() >= 1e-12 {
all_binary = false;
}
parsed.push(Some(v));
trimmed.push(raw);
continue;
}
all_numeric = false;
all_binary = false;
}
level_index.entry(raw.to_string()).or_insert_with(|| {
let idx = levels.len();
levels.push(raw.to_string());
idx
});
parsed.push(None);
trimmed.push(raw);
}
let numeric_column = all_numeric && saw_numeric;
if !numeric_column {
for &i in &missing_positions {
let raw = trimmed[i];
level_index.entry(raw.to_string()).or_insert_with(|| {
let idx = levels.len();
levels.push(raw.to_string());
idx
});
parsed[i] = None;
}
}
let kind = if numeric_column {
if all_binary {
ColumnKindTag::Binary
} else {
ColumnKindTag::Continuous
}
} else {
ColumnKindTag::Categorical
};
if matches!(kind, ColumnKindTag::Categorical) {
sort_levels_canonical(&mut levels);
}
let schema = SchemaColumn {
name: name.to_string(),
kind,
levels: if matches!(kind, ColumnKindTag::Categorical) {
levels
} else {
Vec::new()
},
};
let level_map = if matches!(kind, ColumnKindTag::Categorical) {
Some(
schema
.levels
.iter()
.enumerate()
.map(|(idx, v)| (v.as_str(), idx as f64))
.collect::<HashMap<_, _>>(),
)
} else {
None
};
let mut values = Vec::<f64>::with_capacity(trimmed.len());
for (i, raw) in trimmed.iter().enumerate() {
let raw = *raw;
let val = match kind {
ColumnKindTag::Continuous => parsed[i].ok_or_else(|| {
String::from(DataError::EncodingFailure {
reason: format!(
"internal: continuous column '{}' lost its parsed value at row {} (col {})",
name,
i + 1,
col_index
),
})
})?,
ColumnKindTag::Binary => {
let v = parsed[i].ok_or_else(|| {
String::from(DataError::EncodingFailure {
reason: format!(
"internal: binary column '{}' lost its parsed value at row {} (col {})",
name,
i + 1,
col_index
),
})
})?;
if v.is_finite() && (v - 0.0).abs() >= 1e-12 && (v - 1.0).abs() >= 1e-12 {
return Err(DataError::SchemaMismatch {
reason: format!(
"column '{}' is binary in schema but row {} has value {}; expected 0 or 1",
name,
i + 1,
v
),
}
.into());
}
v
}
ColumnKindTag::Categorical => {
let map = level_map.as_ref().ok_or_else(|| {
String::from(DataError::EncodingFailure {
reason: "internal categorical schema map missing".to_string(),
})
})?;
*map.get(raw).ok_or_else(|| {
String::from(DataError::EncodingFailure {
reason: format!(
"internal: level '{}' missing from freshly built map for column '{}' (col {})",
raw, name, col_index
),
})
})?
}
};
if !val.is_finite() && !is_missing_marker(raw) {
return Err(DataError::InvalidValue {
reason: format!("non-finite value at row {}, column '{}'", i + 1, name),
}
.into());
}
values.push(val);
}
Ok((schema, values))
}
#[cfg(test)]
mod missing_value_inference_tests {
use super::*;
fn rows(cells: &[&[&str]]) -> Vec<StringRecord> {
cells
.iter()
.map(|r| StringRecord::from(r.to_vec()))
.collect()
}
#[test]
fn dtype_categorical_missing_cell_is_not_promoted_to_a_level() {
let column = [Some("g10"), None, Some("g2"), Some("g10")];
let (schema, values) = encode_optional_categorical_column("group", &column)
.expect("encode typed categorical values with a missing cell");
assert_eq!(schema.kind, ColumnKindTag::Categorical);
assert_eq!(schema.levels, vec!["g2", "g10"]);
assert_eq!(values[0], 1.0);
assert!(values[1].is_nan());
assert_eq!(values[2], 0.0);
assert_eq!(values[3], 1.0);
}
#[test]
fn a_numeric_column_containing_na_can_never_infer_categorical() {
let ds = encode_recordswith_inferred_schema(
vec!["parker".to_string()],
rows(&[&["94.4"], &["NA"], &["65.0"], &["88.0"], &["NA"]]),
)
.expect("encode a numeric column carrying NA");
assert_eq!(
ds.schema.columns[0].kind,
ColumnKindTag::Continuous,
"a column whose present cells are all numeric is numeric-with-missing, \
never a factor over its own measurements"
);
assert!(
ds.schema.columns[0].levels.is_empty(),
"measurements must not be recorded as factor levels"
);
let col: Vec<f64> = ds.values.column(0).to_vec();
assert_eq!(col[0], 94.4);
assert_eq!(col[2], 65.0);
assert_eq!(col[3], 88.0);
assert!(col[1].is_nan() && col[4].is_nan(), "NA must encode as NaN");
assert_eq!(
col.iter().filter(|v| v.is_finite()).count(),
3,
"is_finite() must count exactly the present cells"
);
}
#[test]
fn na_stays_a_level_in_a_genuinely_categorical_column() {
let ds = encode_recordswith_inferred_schema(
vec!["country".to_string()],
rows(&[&["NA"], &["ZA"], &["BW"], &["NA"]]),
)
.expect("encode a categorical column whose labels include NA");
assert_eq!(ds.schema.columns[0].kind, ColumnKindTag::Categorical);
assert!(
ds.schema.columns[0].levels.iter().any(|l| l == "NA"),
"NA is a country here, not missingness: levels were {:?}",
ds.schema.columns[0].levels
);
assert!(
ds.values.column(0).iter().all(|v| v.is_finite()),
"a categorical column carries level codes, never NaN"
);
}
#[test]
fn a_binary_column_containing_na_stays_binary_with_nan_holes() {
let ds = encode_recordswith_inferred_schema(
vec!["event".to_string()],
rows(&[&["1"], &["NA"], &["0"], &["1"]]),
)
.expect("encode a binary column carrying NA");
assert_eq!(ds.schema.columns[0].kind, ColumnKindTag::Binary);
let col: Vec<f64> = ds.values.column(0).to_vec();
assert_eq!((col[0], col[2], col[3]), (1.0, 0.0, 1.0));
assert!(col[1].is_nan());
}
#[test]
fn a_parsed_non_finite_literal_still_fails_loudly() {
let err = encode_recordswith_inferred_schema(
vec!["x".to_string()],
rows(&[&["1.0"], &["inf"], &["2.0"]]),
)
.expect_err("a literal infinity is a data error, not a missing value");
assert!(
err.contains("non-finite"),
"expected the non-finite guard, got: {err}"
);
}
}
#[cfg(test)]
mod tests {
use super::*;
use arrow::array::ArrayRef;
use arrow::datatypes::{Field, Schema};
use arrow::error::ArrowError;
use arrow::record_batch::{RecordBatch, RecordBatchIterator};
use std::sync::Arc;
fn encode_single_arrow_array(array: ArrayRef) -> Result<EncodedDataset, DataError> {
let schema = Arc::new(Schema::new(vec![Field::new(
"source",
array.data_type().clone(),
true,
)]));
let batch = RecordBatch::try_new(schema.clone(), vec![array]).expect("record batch");
let batches: Vec<Result<RecordBatch, ArrowError>> = vec![Ok(batch)];
let mut reader = RecordBatchIterator::new(batches, schema);
encode_arrow_record_batch_reader_with_inferred_schema(
&mut reader,
vec!["normalized".to_string()],
)
}
#[test]
fn arrow_reader_streams_typed_columns_in_supplied_order() {
use arrow::array::{
Array, BooleanArray, DictionaryArray, Float32Array, Int8Array, Int32Array, Int64Array,
LargeStringArray, StringArray,
};
use arrow::datatypes::Int8Type;
let string_dictionary_1: DictionaryArray<Int8Type> =
vec!["beta", "alpha"].into_iter().collect();
let string_dictionary_2: DictionaryArray<Int8Type> = vec!["beta"].into_iter().collect();
let numeric_dictionary_1 = DictionaryArray::<Int8Type>::new(
Int8Array::from(vec![0, 1]),
Arc::new(Int64Array::from(vec![5, 7])),
);
let numeric_dictionary_2 = DictionaryArray::<Int8Type>::new(
Int8Array::from(vec![0]),
Arc::new(Int64Array::from(vec![5])),
);
let schema = Arc::new(Schema::new(vec![
Field::new("source_float", arrow::datatypes::DataType::Float32, false),
Field::new("source_integer", arrow::datatypes::DataType::Int32, false),
Field::new("source_flag", arrow::datatypes::DataType::Boolean, false),
Field::new("source_utf8", arrow::datatypes::DataType::Utf8, false),
Field::new(
"source_large_utf8",
arrow::datatypes::DataType::LargeUtf8,
false,
),
Field::new(
"source_dictionary_string",
string_dictionary_1.data_type().clone(),
false,
),
Field::new(
"source_dictionary_number",
numeric_dictionary_1.data_type().clone(),
false,
),
]));
let batch_1 = RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(Float32Array::from(vec![1.5, 2.5])) as ArrayRef,
Arc::new(Int32Array::from(vec![0, 1])),
Arc::new(BooleanArray::from(vec![true, false])),
Arc::new(StringArray::from(vec!["item10", "item2"])),
Arc::new(LargeStringArray::from(vec!["z", "a"])),
Arc::new(string_dictionary_1),
Arc::new(numeric_dictionary_1),
],
)
.expect("first record batch");
let batch_2 = RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(Float32Array::from(vec![-4.0])) as ArrayRef,
Arc::new(Int32Array::from(vec![1])),
Arc::new(BooleanArray::from(vec![true])),
Arc::new(StringArray::from(vec!["item1"])),
Arc::new(LargeStringArray::from(vec!["z"])),
Arc::new(string_dictionary_2),
Arc::new(numeric_dictionary_2),
],
)
.expect("second record batch");
let batches: Vec<Result<RecordBatch, ArrowError>> = vec![Ok(batch_1), Ok(batch_2)];
let mut reader = RecordBatchIterator::new(batches, schema);
let headers = [
"float",
"integer",
"flag",
"utf8",
"large_utf8",
"dictionary_string",
"dictionary_number",
]
.map(str::to_string)
.to_vec();
let dataset =
encode_arrow_record_batch_reader_with_inferred_schema(&mut reader, headers.clone())
.expect("Arrow stream should encode");
assert_eq!(dataset.headers, headers);
assert_eq!(
dataset.column_kinds,
vec![
ColumnKindTag::Continuous,
ColumnKindTag::Binary,
ColumnKindTag::Binary,
ColumnKindTag::Categorical,
ColumnKindTag::Categorical,
ColumnKindTag::Categorical,
ColumnKindTag::Continuous,
]
);
assert_eq!(
dataset.values,
ndarray::arr2(&[
[1.5, 0.0, 1.0, 2.0, 1.0, 1.0, 5.0],
[2.5, 1.0, 0.0, 1.0, 0.0, 0.0, 7.0],
[-4.0, 1.0, 1.0, 0.0, 1.0, 1.0, 5.0],
])
);
assert_eq!(
dataset.schema.columns[3].levels,
vec!["item1", "item2", "item10"]
);
assert_eq!(dataset.schema.columns[4].levels, vec!["a", "z"]);
assert_eq!(dataset.schema.columns[5].levels, vec!["alpha", "beta"]);
assert!(
dataset
.schema
.columns
.iter()
.zip(dataset.headers.iter())
.all(|(column, header)| column.name == *header)
);
}
#[test]
fn arrow_reader_rejects_empty_duplicate_and_mismatched_headers() {
let schema = Arc::new(Schema::new(vec![Field::new(
"source",
arrow::datatypes::DataType::Int32,
false,
)]));
let mut empty_name_reader = RecordBatchIterator::new(
Vec::<Result<RecordBatch, ArrowError>>::new(),
schema.clone(),
);
let empty_name = encode_arrow_record_batch_reader_with_inferred_schema(
&mut empty_name_reader,
vec![" ".to_string()],
)
.expect_err("blank header should fail");
assert!(matches!(empty_name, DataError::EmptyInput { .. }));
let mut duplicate_reader = RecordBatchIterator::new(
Vec::<Result<RecordBatch, ArrowError>>::new(),
Arc::new(Schema::new(vec![
Field::new("a", arrow::datatypes::DataType::Int32, false),
Field::new("b", arrow::datatypes::DataType::Int32, false),
])),
);
let duplicate = encode_arrow_record_batch_reader_with_inferred_schema(
&mut duplicate_reader,
vec!["x".to_string(), "x".to_string()],
)
.expect_err("duplicate header should fail");
assert!(matches!(duplicate, DataError::SchemaMismatch { .. }));
let mut mismatch_reader =
RecordBatchIterator::new(Vec::<Result<RecordBatch, ArrowError>>::new(), schema);
let mismatch = encode_arrow_record_batch_reader_with_inferred_schema(
&mut mismatch_reader,
vec!["x".to_string(), "y".to_string()],
)
.expect_err("header count mismatch should fail");
assert!(matches!(mismatch, DataError::SchemaMismatch { .. }));
}
#[test]
fn arrow_reader_preserves_missing_cells_and_rejects_unsupported_types() {
use arrow::array::{Date32Array, DictionaryArray, Float64Array, Int8Array, StringArray};
use arrow::datatypes::Int8Type;
let null_numeric =
encode_single_arrow_array(Arc::new(Float64Array::from(vec![Some(1.0), None])))
.expect("numeric null should remain representable until model projection");
assert_eq!(null_numeric.values[[0, 0]], 1.0);
assert!(null_numeric.values[[1, 0]].is_nan());
let null_dictionary = DictionaryArray::<Int8Type>::new(
Int8Array::from(vec![0, 1, 2]),
Arc::new(StringArray::from(vec![Some("present"), Some(""), None])),
);
let logical_null = encode_single_arrow_array(Arc::new(null_dictionary))
.expect("empty and null dictionary values should remain representable");
assert_eq!(logical_null.schema.columns[0].levels, vec!["present"]);
assert_eq!(logical_null.values[[0, 0]], 0.0);
assert!(logical_null.values[[1, 0]].is_nan());
assert!(logical_null.values[[2, 0]].is_nan());
let nonfinite = encode_single_arrow_array(Arc::new(Float64Array::from(vec![
f64::NAN,
f64::INFINITY,
f64::NEG_INFINITY,
])))
.expect("non-finite values should remain representable until model projection");
assert!(nonfinite.values.column(0).iter().all(|value| value.is_nan()));
assert_eq!(
nonfinite.column_kinds,
vec![ColumnKindTag::Continuous],
"an all-missing typed numeric column is not vacuously binary"
);
let unsupported = encode_single_arrow_array(Arc::new(Date32Array::from(vec![1])))
.expect_err("date column should fail");
assert!(matches!(&unsupported, DataError::InvalidValue { .. }));
assert!(
unsupported
.to_string()
.contains("unsupported Arrow column type")
);
}
#[test]
fn encode_records_rejects_empty_input() {
let headers = vec!["x".to_string()];
let schema = DataSchema {
columns: vec![SchemaColumn {
name: "x".to_string(),
kind: ColumnKindTag::Continuous,
levels: Vec::new(),
}],
};
let err = encode_recordswith_inferred_schema(headers.clone(), Vec::new())
.expect_err("empty inferred records should error");
assert_eq!(err, "table data cannot be empty");
let err =
encode_recordswith_schema(headers, Vec::new(), &schema, UnseenCategoryPolicy::Error)
.expect_err("empty schema-guided records should error");
assert_eq!(err, "table data cannot be empty");
}
#[test]
fn column_major_matches_record_driven_inferred_encode() {
let headers = vec!["cont".to_string(), "bin".to_string(), "cat".to_string()];
let raw_rows = vec![
vec!["1.5", "0", "a"],
vec!["2.0", "1", "b"],
vec!["-3.25", "1", "a"],
vec!["0.0", "0", "c"],
];
let records: Vec<StringRecord> = raw_rows
.iter()
.map(|r| StringRecord::from(r.clone()))
.collect();
let record_ds = encode_recordswith_inferred_schema(headers.clone(), records)
.expect("record-driven encode");
for (j, name) in headers.iter().enumerate() {
let column: Vec<&str> = raw_rows.iter().map(|r| r[j]).collect();
let (schema_col, values) =
infer_and_encode_column_major(name, &column, j + 1).expect("column-major encode");
assert_eq!(schema_col.kind, record_ds.schema.columns[j].kind);
assert_eq!(schema_col.levels, record_ds.schema.columns[j].levels);
for (i, v) in values.iter().enumerate() {
assert_eq!(*v, record_ds.values[[i, j]], "row {i} col {name}");
}
}
}
#[test]
fn encode_records_can_encode_unseen_named_categorical_column() {
let schema = DataSchema {
columns: vec![
SchemaColumn {
name: "g".to_string(),
kind: ColumnKindTag::Categorical,
levels: vec!["a".to_string(), "b".to_string()],
},
SchemaColumn {
name: "x".to_string(),
kind: ColumnKindTag::Categorical,
levels: vec!["low".to_string(), "high".to_string()],
},
],
};
let headers = vec!["g".to_string(), "x".to_string()];
let records = vec![StringRecord::from(vec!["new-group", "low"])];
let policy =
UnseenCategoryPolicy::encode_unknown_for_columns(HashSet::from(["g".to_string()]));
let ds =
encode_recordswith_schema(headers, records, &schema, policy).expect("encoded dataset");
assert_eq!(ds.values[[0, 0]], 2.0);
assert_eq!(ds.values[[0, 1]], 0.0);
}
#[test]
fn categorical_encoder_consumes_labels_and_remaps_canonically() {
use ndarray::Array1;
let mut encoder = CategoricalEncoder::default();
let mut encoded = Array1::from_vec(
["item10", "item2", "item1", "item2"]
.into_iter()
.map(|label| encoder.encode(label) as f64)
.collect(),
);
let levels = encoder.finish(encoded.view_mut(), LevelOrder::Canonical);
assert_eq!(levels, vec!["item1", "item2", "item10"]);
assert_eq!(encoded.to_vec(), vec![2.0, 1.0, 0.0, 1.0]);
}
#[test]
fn complete_delimited_schema_encodes_projected_rows_directly() {
let dir = tempfile::tempdir().expect("tempdir");
let path = dir.path().join("schema_direct.csv");
std::fs::write(
&path,
"y,group,flag,unused\n1.5,b,0,first\n2.5,a,1,second\n",
)
.expect("write csv");
let schema = DataSchema {
columns: vec![
SchemaColumn {
name: "group".to_string(),
kind: ColumnKindTag::Categorical,
levels: vec!["a".to_string(), "b".to_string()],
},
SchemaColumn {
name: "flag".to_string(),
kind: ColumnKindTag::Binary,
levels: Vec::new(),
},
SchemaColumn {
name: "y".to_string(),
kind: ColumnKindTag::Continuous,
levels: Vec::new(),
},
],
};
let loaded = load_datasetwith_schema_projected(
&path,
&schema,
UnseenCategoryPolicy::Error,
&["y".to_string(), "group".to_string(), "flag".to_string()],
)
.expect("schema-guided projected load");
assert_eq!(loaded.headers, vec!["y", "group", "flag"]);
assert_eq!(loaded.values.row(0).to_vec(), vec![1.5, 1.0, 0.0]);
assert_eq!(loaded.values.row(1).to_vec(), vec![2.5, 0.0, 1.0]);
assert_eq!(
loaded.column_kinds,
vec![
ColumnKindTag::Continuous,
ColumnKindTag::Categorical,
ColumnKindTag::Binary,
]
);
}
#[test]
fn direct_parquet_decoder_preserves_string_encounter_order() {
use arrow::array::DictionaryArray;
use arrow::datatypes::Int8Type;
use ndarray::Array1;
let dictionary: DictionaryArray<Int8Type> =
vec!["beta", "alpha", "beta"].into_iter().collect();
let mut encoded = Array1::<f64>::zeros(dictionary.len());
let mut encoder = CategoricalEncoder::default();
let mut all_binary = true;
let mut saw_numeric = false;
decode_arrow_batch_column_into(
&dictionary,
0,
"group",
true,
encoded.view_mut(),
Some(&mut encoder),
&mut all_binary,
&mut saw_numeric,
)
.expect("dictionary strings decode directly");
let levels = encoder.finish(encoded.view_mut(), LevelOrder::Encounter);
assert_eq!(levels, vec!["beta", "alpha"]);
assert_eq!(encoded.to_vec(), vec![0.0, 1.0, 0.0]);
}
#[test]
fn numeric_valued_dictionary_column_classifies_and_decodes_as_numeric() {
use arrow::array::{Array, ArrayRef, DictionaryArray, Int8Array, Int64Array};
use arrow::datatypes::{DataType, Int8Type};
use std::sync::Arc;
let keys = Int8Array::from(vec![0i8, 1, 0, 1, 0]);
let dict_values: ArrayRef = Arc::new(Int64Array::from(vec![5i64, 7]));
let dict: DictionaryArray<Int8Type> = DictionaryArray::new(keys, dict_values);
assert!(matches!(dict.data_type(), DataType::Dictionary(_, _)));
assert!(
!arrow_field_is_string(dict.data_type()),
"Dictionary(Int8, Int64) must not be treated as a string column"
);
let str_dict: DictionaryArray<Int8Type> = vec!["a", "b", "a"].into_iter().collect();
assert!(
arrow_field_is_string(str_dict.data_type()),
"Dictionary(Int8, Utf8) must remain a string column"
);
let mut decoded = ndarray::Array1::<f64>::zeros(dict.len());
let mut all_binary = true;
let mut saw_numeric = false;
decode_arrow_batch_column_into(
&dict,
0,
"x",
false,
decoded.view_mut(),
None,
&mut all_binary,
&mut saw_numeric,
)
.expect("numeric dictionary column should decode as numeric");
assert_eq!(decoded.to_vec(), vec![5.0, 7.0, 5.0, 7.0, 5.0]);
assert!(!all_binary);
use arrow::datatypes::{Field, Schema};
use arrow::record_batch::RecordBatch;
use parquet::arrow::ArrowWriter;
let arrow_schema = Arc::new(Schema::new(vec![Field::new(
"x",
dict.data_type().clone(),
false,
)]));
let batch = RecordBatch::try_new(arrow_schema.clone(), vec![Arc::new(dict.clone())])
.expect("record batch with a dictionary numeric column");
let dir = tempfile::tempdir().expect("tempdir");
let path = dir.path().join("dict_numeric.parquet");
{
let file = std::fs::File::create(&path).expect("create parquet");
let mut writer =
ArrowWriter::try_new(file, arrow_schema, None).expect("arrow parquet writer");
writer.write(&batch).expect("write batch");
writer.close().expect("close writer");
}
let inferred =
load_parquet_inferred(&path, &[], &HashSet::new()).expect("inferred parquet load");
assert_eq!(inferred.column_kinds, vec![ColumnKindTag::Continuous]);
assert_eq!(
inferred.values.column(0).to_vec(),
vec![5.0, 7.0, 5.0, 7.0, 5.0]
);
let schema = DataSchema {
columns: vec![SchemaColumn {
name: "x".to_string(),
kind: ColumnKindTag::Continuous,
levels: Vec::new(),
}],
};
let schema_loaded =
load_parquet_with_schema(&path, &schema, UnseenCategoryPolicy::Error, &[])
.expect("dictionary-encoded numeric parquet must load against a Continuous schema");
assert_eq!(schema_loaded.column_kinds, vec![ColumnKindTag::Continuous]);
assert_eq!(
schema_loaded.values.column(0).to_vec(),
vec![5.0, 7.0, 5.0, 7.0, 5.0]
);
}
#[test]
fn encode_records_keeps_unlisted_categorical_columns_strict() {
let schema = DataSchema {
columns: vec![
SchemaColumn {
name: "g".to_string(),
kind: ColumnKindTag::Categorical,
levels: vec!["a".to_string(), "b".to_string()],
},
SchemaColumn {
name: "x".to_string(),
kind: ColumnKindTag::Categorical,
levels: vec!["low".to_string(), "high".to_string()],
},
],
};
let headers = vec!["g".to_string(), "x".to_string()];
let records = vec![StringRecord::from(vec!["a", "new-level"])];
let policy =
UnseenCategoryPolicy::encode_unknown_for_columns(HashSet::from(["g".to_string()]));
let err = encode_recordswith_schema(headers, records, &schema, policy)
.expect_err("ordinary categorical column should stay strict");
assert!(err.contains("unseen level 'new-level' in categorical column 'x'"));
}
#[test]
fn sentinel_strip_present_returns_rest_and_true() {
let marked = format!("{}{}", CATEGORICAL_CELL_SENTINEL, "hello");
let (rest, found) = strip_categorical_sentinel(&marked);
assert_eq!(rest, "hello");
assert!(found);
}
#[test]
fn sentinel_strip_absent_returns_original_and_false() {
let (rest, found) = strip_categorical_sentinel("hello");
assert_eq!(rest, "hello");
assert!(!found);
}
#[test]
fn sentinel_strip_empty_string_returns_empty_and_false() {
let (rest, found) = strip_categorical_sentinel("");
assert_eq!(rest, "");
assert!(!found);
}
#[test]
fn sentinel_strip_only_sentinel_returns_empty_and_true() {
let marked = CATEGORICAL_CELL_SENTINEL.to_string();
let (rest, found) = strip_categorical_sentinel(&marked);
assert_eq!(rest, "");
assert!(found);
}
#[test]
fn feature_ranges_two_columns() {
let values = ndarray::arr2(&[[1.0_f64, 10.0], [3.0, 20.0], [2.0, 15.0]]);
let ds = EncodedDataset {
headers: vec!["a".to_string(), "b".to_string()],
values,
schema: DataSchema { columns: vec![] },
column_kinds: vec![ColumnKindTag::Continuous, ColumnKindTag::Continuous],
};
let ranges = ds.feature_ranges();
assert_eq!(ranges.len(), 2);
assert_eq!(ranges[0], (1.0, 3.0));
assert_eq!(ranges[1], (10.0, 20.0));
}
#[test]
fn feature_ranges_single_row_min_equals_max() {
let values = ndarray::arr2(&[[5.0_f64, -3.0]]);
let ds = EncodedDataset {
headers: vec!["x".to_string(), "y".to_string()],
values,
schema: DataSchema { columns: vec![] },
column_kinds: vec![ColumnKindTag::Continuous, ColumnKindTag::Continuous],
};
let ranges = ds.feature_ranges();
assert_eq!(ranges[0], (5.0, 5.0));
assert_eq!(ranges[1], (-3.0, -3.0));
}
#[test]
fn feature_ranges_all_nan_defaults_to_zero() {
let values = ndarray::arr2(&[[f64::NAN], [f64::NAN]]);
let ds = EncodedDataset {
headers: vec!["x".to_string()],
values,
schema: DataSchema { columns: vec![] },
column_kinds: vec![ColumnKindTag::Continuous],
};
let ranges = ds.feature_ranges();
assert_eq!(ranges[0], (0.0, 0.0));
}
#[test]
fn column_map_indexes_by_name() {
let values = ndarray::arr2(&[[0.0_f64, 1.0], [2.0, 3.0]]);
let ds = EncodedDataset {
headers: vec!["alpha".to_string(), "beta".to_string()],
values,
schema: DataSchema { columns: vec![] },
column_kinds: vec![ColumnKindTag::Continuous, ColumnKindTag::Continuous],
};
let map = ds.column_map();
assert_eq!(map["alpha"], 0);
assert_eq!(map["beta"], 1);
assert_eq!(map.len(), 2);
}
#[test]
fn shared_prefix_identical_strings() {
assert_eq!(shared_prefix("hello", "hello"), 5);
}
#[test]
fn shared_prefix_no_common_prefix() {
assert_eq!(shared_prefix("abc", "xyz"), 0);
}
#[test]
fn shared_prefix_partial_match() {
assert_eq!(shared_prefix("foobar", "foobaz"), 5);
}
#[test]
fn shared_prefix_one_empty() {
assert_eq!(shared_prefix("", "hello"), 0);
assert_eq!(shared_prefix("hello", ""), 0);
}
#[test]
fn shared_prefix_both_empty() {
assert_eq!(shared_prefix("", ""), 0);
}
#[test]
fn shared_prefix_shorter_string_is_prefix() {
assert_eq!(shared_prefix("foo", "foobar"), 3);
}
#[test]
fn detect_format_csv() {
let path = std::path::Path::new("data.csv");
assert_eq!(detect_format(path).unwrap(), DataFormat::Csv);
}
#[test]
fn detect_format_tsv() {
assert_eq!(
detect_format(std::path::Path::new("data.tsv")).unwrap(),
DataFormat::Tsv
);
assert_eq!(
detect_format(std::path::Path::new("data.txt")).unwrap(),
DataFormat::Tsv
);
assert_eq!(
detect_format(std::path::Path::new("data.tab")).unwrap(),
DataFormat::Tsv
);
}
#[test]
fn detect_format_parquet() {
assert_eq!(
detect_format(std::path::Path::new("data.parquet")).unwrap(),
DataFormat::Parquet
);
assert_eq!(
detect_format(std::path::Path::new("data.pq")).unwrap(),
DataFormat::Parquet
);
assert_eq!(
detect_format(std::path::Path::new("data.pqt")).unwrap(),
DataFormat::Parquet
);
}
#[test]
fn detect_format_uppercase_extension() {
assert_eq!(
detect_format(std::path::Path::new("data.CSV")).unwrap(),
DataFormat::Csv
);
}
#[test]
fn detect_format_unknown_extension_is_error() {
let err = detect_format(std::path::Path::new("data.json")).unwrap_err();
let msg = format!("{err:?}");
assert!(
msg.contains("json") || msg.contains("unsupported"),
"error should mention extension, got: {msg}"
);
}
#[test]
fn strip_categorical_sentinel_marked_cell() {
let marked = "\u{0}hello";
let (text, found) = strip_categorical_sentinel(marked);
assert!(found);
assert_eq!(text, "hello");
}
#[test]
fn strip_categorical_sentinel_unmarked_cell() {
let (text, found) = strip_categorical_sentinel("plain");
assert!(!found);
assert_eq!(text, "plain");
}
#[test]
fn strip_categorical_sentinel_empty_string() {
let (text, found) = strip_categorical_sentinel("");
assert!(!found);
assert_eq!(text, "");
}
#[test]
fn strip_categorical_sentinel_only_sentinel() {
let s = "\u{0}";
let (text, found) = strip_categorical_sentinel(s);
assert!(found);
assert_eq!(text, "");
}
#[test]
fn projected_headers_selects_by_index() {
let all = vec![
"a".to_string(),
"b".to_string(),
"c".to_string(),
"d".to_string(),
];
let selected = projected_headers(&all, &[1, 3]);
assert_eq!(selected, vec!["b".to_string(), "d".to_string()]);
}
#[test]
fn projected_headers_empty_selection() {
let all = vec!["x".to_string(), "y".to_string()];
let selected = projected_headers(&all, &[]);
assert!(selected.is_empty());
}
#[test]
fn projected_headers_all_indices() {
let all = vec!["p".to_string(), "q".to_string()];
let selected = projected_headers(&all, &[0, 1]);
assert_eq!(selected, all);
}
#[test]
fn canonical_level_bits_collapses_signed_zero() {
let pos = 0.0_f64;
let neg = -0.0_f64;
assert_ne!(
pos.to_bits(),
neg.to_bits(),
"precondition: raw bits differ"
);
assert_eq!(pos, neg, "precondition: numerically equal");
assert_eq!(canonical_level_bits(pos), canonical_level_bits(neg));
assert_eq!(canonical_level_bits(neg), 0.0_f64.to_bits());
assert_eq!(canonical_level_bits(-1.0 * 0.0), 0.0_f64.to_bits());
assert_eq!(canonical_level_bits(0.0 - 0.0), 0.0_f64.to_bits());
}
#[test]
fn canonical_level_bits_is_bit_stable_on_ordinary_values() {
for &v in &[
1.0_f64,
-1.0,
2.5,
-3.75,
1e300,
-1e-300,
f64::MIN,
f64::MAX,
] {
assert_eq!(canonical_level_bits(v), v.to_bits(), "value {v}");
}
assert_ne!(canonical_level_bits(1.0), canonical_level_bits(2.0));
assert_ne!(canonical_level_bits(0.0), canonical_level_bits(1.0));
assert_ne!(
canonical_level_bits(f64::INFINITY),
canonical_level_bits(f64::NEG_INFINITY)
);
}
#[test]
fn canonical_level_bits_collapses_nan_payloads() {
let a = f64::NAN;
let b = f64::from_bits(0x7ff8_0000_0000_0001); let c = -f64::NAN; assert!(a.is_nan() && b.is_nan() && c.is_nan());
assert_eq!(canonical_level_bits(a), canonical_level_bits(b));
assert_eq!(canonical_level_bits(a), canonical_level_bits(c));
}
#[test]
fn canonical_level_bits_is_idempotent() {
for &v in &[0.0_f64, -0.0, 1.0, -2.0, f64::NAN] {
let once = canonical_level_bits(v);
let twice = canonical_level_bits(f64::from_bits(once));
assert_eq!(once, twice, "value {v}");
}
}
#[test]
fn fit_boundary_reports_each_degenerate_column_by_name() {
let cases = [
(
vec![0.0, f64::NAN, 1.0],
"has non-finite value NaN at row 2",
),
(
vec![0.0, f64::INFINITY, 1.0],
"has non-finite value inf at row 2",
),
(
vec![0.0, f64::NEG_INFINITY, 1.0],
"has non-finite value -inf at row 2",
),
(
vec![f64::NAN, 2.0, f64::NAN],
"has only one non-missing value",
),
];
for (values, expected) in cases {
let dataset = EncodedDataset {
headers: vec!["temperature".to_string()],
values: Array2::from_shape_vec((3, 1), values).unwrap(),
schema: DataSchema {
columns: vec![SchemaColumn {
name: "temperature".to_string(),
kind: ColumnKindTag::Continuous,
levels: Vec::new(),
}],
},
column_kinds: vec![ColumnKindTag::Continuous],
};
let error = dataset.validate_fit_boundary().unwrap_err();
assert!(matches!(error, DataError::DegenerateColumn { .. }));
assert_eq!(
error.to_string(),
format!("column 'temperature' {expected}")
);
}
}
#[test]
fn fit_boundary_rejects_empty_duplicate_and_one_level_factor() {
let cases = [
EncodedDataset {
headers: vec!["x".into()],
values: Array2::zeros((0, 1)),
schema: DataSchema {
columns: vec![SchemaColumn {
name: "x".into(),
kind: ColumnKindTag::Continuous,
levels: vec![],
}],
},
column_kinds: vec![ColumnKindTag::Continuous],
},
EncodedDataset {
headers: vec!["x".into(), "x".into()],
values: Array2::from_shape_vec((2, 2), vec![0.0, 1.0, 1.0, 0.0]).unwrap(),
schema: DataSchema { columns: vec![] },
column_kinds: vec![ColumnKindTag::Continuous; 2],
},
EncodedDataset {
headers: vec!["group".into()],
values: Array2::zeros((2, 1)),
schema: DataSchema {
columns: vec![SchemaColumn {
name: "group".into(),
kind: ColumnKindTag::Categorical,
levels: vec!["only".into()],
}],
},
column_kinds: vec![ColumnKindTag::Categorical],
},
];
let expected = [
"column '<table>' has no observations",
"column 'x' has a duplicate name",
"column 'group' is a factor with fewer than two levels",
];
for (dataset, expected) in cases.into_iter().zip(expected) {
assert_eq!(
dataset.validate_fit_boundary().unwrap_err().to_string(),
expected
);
}
}
}