use std::path::Path;
use std::sync::Arc;
use arrow::array::{Float64Builder, Int32Builder, RecordBatch};
use arrow::datatypes::{DataType, Field, Schema};
use crate::extensions::FphaHyperplaneRow;
use crate::output::atomic::{write_json_atomic, write_parquet_atomic};
use crate::output::error::OutputError;
use crate::output::parquet_config::ParquetWriterConfig;
use crate::output::stochastic::ensure_parent_dir;
pub fn write_fpha_hyperplanes(path: &Path, rows: &[FphaHyperplaneRow]) -> Result<(), OutputError> {
ensure_parent_dir(path)?;
let config = ParquetWriterConfig::default();
let batch = build_fpha_hyperplanes_batch(rows)?;
write_parquet_atomic(path, &batch, &config)
}
fn fpha_hyperplanes_schema() -> Schema {
Schema::new(vec![
Field::new("hydro_id", DataType::Int32, false),
Field::new("stage_id", DataType::Int32, true),
Field::new("plane_id", DataType::Int32, false),
Field::new("gamma_0", DataType::Float64, false),
Field::new("gamma_v", DataType::Float64, false),
Field::new("gamma_q", DataType::Float64, false),
Field::new("gamma_s", DataType::Float64, false),
Field::new("kappa", DataType::Float64, true),
Field::new("valid_v_min_hm3", DataType::Float64, true),
Field::new("valid_v_max_hm3", DataType::Float64, true),
Field::new("valid_q_max_m3s", DataType::Float64, true),
])
}
#[allow(clippy::similar_names)]
fn build_fpha_hyperplanes_batch(rows: &[FphaHyperplaneRow]) -> Result<RecordBatch, OutputError> {
let n = rows.len();
let mut hydro_id_col = Int32Builder::with_capacity(n);
let mut stage_id_col = Int32Builder::with_capacity(n);
let mut plane_id_col = Int32Builder::with_capacity(n);
let mut gamma_0_col = Float64Builder::with_capacity(n);
let mut gamma_v_col = Float64Builder::with_capacity(n);
let mut gamma_q_col = Float64Builder::with_capacity(n);
let mut gamma_s_col = Float64Builder::with_capacity(n);
let mut kappa_col = Float64Builder::with_capacity(n);
let mut valid_v_min_col = Float64Builder::with_capacity(n);
let mut valid_v_max_col = Float64Builder::with_capacity(n);
let mut valid_q_max_col = Float64Builder::with_capacity(n);
for row in rows {
hydro_id_col.append_value(row.hydro_id.0);
stage_id_col.append_option(row.stage_id);
plane_id_col.append_value(row.plane_id);
gamma_0_col.append_value(row.gamma_0);
gamma_v_col.append_value(row.gamma_v);
gamma_q_col.append_value(row.gamma_q);
gamma_s_col.append_value(row.gamma_s);
kappa_col.append_value(row.kappa);
valid_v_min_col.append_option(row.valid_v_min_hm3);
valid_v_max_col.append_option(row.valid_v_max_hm3);
valid_q_max_col.append_option(row.valid_q_max_m3s);
}
RecordBatch::try_new(
Arc::new(fpha_hyperplanes_schema()),
vec![
Arc::new(hydro_id_col.finish()),
Arc::new(stage_id_col.finish()),
Arc::new(plane_id_col.finish()),
Arc::new(gamma_0_col.finish()),
Arc::new(gamma_v_col.finish()),
Arc::new(gamma_q_col.finish()),
Arc::new(gamma_s_col.finish()),
Arc::new(kappa_col.finish()),
Arc::new(valid_v_min_col.finish()),
Arc::new(valid_v_max_col.finish()),
Arc::new(valid_q_max_col.finish()),
],
)
.map_err(|e| OutputError::serialization("fpha_hyperplanes", e.to_string()))
}
pub fn write_hydro_model_summary(
path: &Path,
summary: &impl serde::Serialize,
) -> Result<(), OutputError> {
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent).map_err(|e| OutputError::io(parent, e))?;
}
write_json_atomic(path, summary, "hydro_models")
}
pub fn read_hydro_model_summary<T: serde::de::DeserializeOwned>(
path: &Path,
) -> Result<T, OutputError> {
let content = std::fs::read_to_string(path).map_err(|e| OutputError::io(path, e))?;
serde_json::from_str(&content).map_err(|e| OutputError::ManifestError {
manifest_type: "hydro_models".to_string(),
message: e.to_string(),
})
}
#[cfg(test)]
#[allow(
clippy::doc_markdown,
clippy::expect_used,
clippy::float_cmp,
clippy::panic,
clippy::too_many_lines,
clippy::unwrap_used
)]
mod tests {
use super::*;
use cobre_core::EntityId;
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
use tempfile::tempdir;
use crate::extensions::parse_fpha_hyperplanes;
fn make_row(hydro_id: i32, plane_id: i32, gamma_0: f64, kappa: f64) -> FphaHyperplaneRow {
FphaHyperplaneRow {
hydro_id: EntityId::from(hydro_id),
stage_id: None,
plane_id,
gamma_0,
gamma_v: 0.0023,
gamma_q: 0.892,
gamma_s: -0.015,
kappa,
valid_v_min_hm3: None,
valid_v_max_hm3: None,
valid_q_max_m3s: None,
}
}
#[test]
fn round_trip_5_rows_hydro_66() {
let rows = vec![
FphaHyperplaneRow {
hydro_id: EntityId::from(66),
stage_id: None,
plane_id: 0,
gamma_0: 1250.5,
gamma_v: 0.0023,
gamma_q: 0.892,
gamma_s: -0.015,
kappa: 0.985,
valid_v_min_hm3: None,
valid_v_max_hm3: None,
valid_q_max_m3s: None,
},
FphaHyperplaneRow {
hydro_id: EntityId::from(66),
stage_id: None,
plane_id: 1,
gamma_0: 1180.2,
gamma_v: 0.0031,
gamma_q: 0.875,
gamma_s: -0.012,
kappa: 0.985,
valid_v_min_hm3: None,
valid_v_max_hm3: None,
valid_q_max_m3s: None,
},
FphaHyperplaneRow {
hydro_id: EntityId::from(66),
stage_id: None,
plane_id: 2,
gamma_0: 1320.8,
gamma_v: 0.0018,
gamma_q: 0.901,
gamma_s: -0.018,
kappa: 0.985,
valid_v_min_hm3: None,
valid_v_max_hm3: None,
valid_q_max_m3s: None,
},
FphaHyperplaneRow {
hydro_id: EntityId::from(66),
stage_id: None,
plane_id: 3,
gamma_0: 1095.4,
gamma_v: 0.0042,
gamma_q: 0.858,
gamma_s: -0.010,
kappa: 0.985,
valid_v_min_hm3: None,
valid_v_max_hm3: None,
valid_q_max_m3s: None,
},
FphaHyperplaneRow {
hydro_id: EntityId::from(66),
stage_id: None,
plane_id: 4,
gamma_0: 1410.1,
gamma_v: 0.0012,
gamma_q: 0.915,
gamma_s: -0.022,
kappa: 0.985,
valid_v_min_hm3: None,
valid_v_max_hm3: None,
valid_q_max_m3s: None,
},
];
let tmp = tempdir().expect("tempdir");
let path = tmp.path().join("fpha_hyperplanes.parquet");
write_fpha_hyperplanes(&path, &rows).expect("write must succeed");
assert!(path.exists(), "file must exist after write");
let parsed = parse_fpha_hyperplanes(&path).expect("parse must succeed");
assert_eq!(parsed.len(), 5, "must have 5 rows");
for (written, read) in rows.iter().zip(parsed.iter()) {
assert_eq!(read.hydro_id, written.hydro_id, "hydro_id mismatch");
assert_eq!(read.plane_id, written.plane_id, "plane_id mismatch");
assert_eq!(read.stage_id, written.stage_id, "stage_id mismatch");
assert!(
(read.gamma_0 - written.gamma_0).abs() < 1e-10,
"gamma_0 mismatch: {} vs {}",
read.gamma_0,
written.gamma_0
);
assert!(
(read.gamma_v - written.gamma_v).abs() < 1e-10,
"gamma_v mismatch"
);
assert!(
(read.gamma_q - written.gamma_q).abs() < 1e-10,
"gamma_q mismatch"
);
assert!(
(read.gamma_s - written.gamma_s).abs() < 1e-10,
"gamma_s mismatch"
);
assert!(
(read.kappa - written.kappa).abs() < 1e-10,
"kappa mismatch: {} vs {}",
read.kappa,
written.kappa
);
}
}
#[test]
fn empty_slice_produces_valid_parquet_with_zero_rows() {
let tmp = tempdir().expect("tempdir");
let path = tmp.path().join("fpha_hyperplanes.parquet");
write_fpha_hyperplanes(&path, &[]).expect("write must succeed for empty slice");
assert!(path.exists(), "file must exist after write");
let parsed = parse_fpha_hyperplanes(&path).expect("parse must succeed");
assert!(parsed.is_empty(), "must have 0 rows for empty input");
}
#[test]
fn schema_has_exactly_11_fields_with_correct_names_and_types() {
let rows = vec![make_row(5, 0, 1000.0, 0.97)];
let tmp = tempdir().expect("tempdir");
let path = tmp.path().join("fpha_hyperplanes.parquet");
write_fpha_hyperplanes(&path, &rows).expect("write must succeed");
let file = std::fs::File::open(&path).unwrap();
let builder = ParquetRecordBatchReaderBuilder::try_new(file).unwrap();
let schema = builder.schema().clone();
assert_eq!(schema.fields().len(), 11, "schema must have 11 fields");
let expected_fields: &[(&str, bool)] = &[
("hydro_id", false),
("stage_id", true),
("plane_id", false),
("gamma_0", false),
("gamma_v", false),
("gamma_q", false),
("gamma_s", false),
("kappa", true),
("valid_v_min_hm3", true),
("valid_v_max_hm3", true),
("valid_q_max_m3s", true),
];
for (i, (expected_name, expected_nullable)) in expected_fields.iter().enumerate() {
let field = &schema.fields()[i];
assert_eq!(
field.name(),
*expected_name,
"field {i} name: expected {expected_name}, got {}",
field.name()
);
assert_eq!(
field.is_nullable(),
*expected_nullable,
"field {i} ({}) nullable: expected {expected_nullable}, got {}",
field.name(),
field.is_nullable()
);
}
}
#[test]
fn nullable_columns_round_trip_as_none() {
let rows = vec![FphaHyperplaneRow {
hydro_id: EntityId::from(10),
stage_id: None,
plane_id: 0,
gamma_0: 500.0,
gamma_v: 0.001,
gamma_q: 0.85,
gamma_s: -0.01,
kappa: 1.0,
valid_v_min_hm3: None,
valid_v_max_hm3: None,
valid_q_max_m3s: None,
}];
let tmp = tempdir().expect("tempdir");
let path = tmp.path().join("fpha_hyperplanes.parquet");
write_fpha_hyperplanes(&path, &rows).expect("write must succeed");
let parsed = parse_fpha_hyperplanes(&path).expect("parse must succeed");
assert_eq!(parsed.len(), 1);
let row = &parsed[0];
assert!(row.stage_id.is_none(), "stage_id must be None");
assert!(
row.valid_v_min_hm3.is_none(),
"valid_v_min_hm3 must be None"
);
assert!(
row.valid_v_max_hm3.is_none(),
"valid_v_max_hm3 must be None"
);
assert!(
row.valid_q_max_m3s.is_none(),
"valid_q_max_m3s must be None"
);
}
#[test]
fn multi_hydro_rows_sorted_by_parse() {
let rows = vec![
make_row(10, 1, 200.0, 0.99),
make_row(10, 0, 210.0, 0.99),
make_row(5, 2, 300.0, 0.95),
make_row(5, 0, 310.0, 0.95),
make_row(5, 1, 305.0, 0.95),
];
let tmp = tempdir().expect("tempdir");
let path = tmp.path().join("fpha_hyperplanes.parquet");
write_fpha_hyperplanes(&path, &rows).expect("write must succeed");
let parsed = parse_fpha_hyperplanes(&path).expect("parse must succeed");
assert_eq!(parsed.len(), 5);
assert_eq!(parsed[0].hydro_id, EntityId::from(5));
assert_eq!(parsed[0].plane_id, 0);
assert_eq!(parsed[1].hydro_id, EntityId::from(5));
assert_eq!(parsed[1].plane_id, 1);
assert_eq!(parsed[2].hydro_id, EntityId::from(5));
assert_eq!(parsed[2].plane_id, 2);
assert_eq!(parsed[3].hydro_id, EntityId::from(10));
assert_eq!(parsed[3].plane_id, 0);
assert_eq!(parsed[4].hydro_id, EntityId::from(10));
assert_eq!(parsed[4].plane_id, 1);
}
#[test]
fn parent_directory_created_automatically() {
let tmp = tempdir().expect("tempdir");
let path = tmp
.path()
.join("output")
.join("hydro_models")
.join("fpha_hyperplanes.parquet");
assert!(
!path.parent().unwrap().exists(),
"parent dir must not exist before write"
);
write_fpha_hyperplanes(&path, &[make_row(1, 0, 100.0, 1.0)])
.expect("write must succeed even when parent dirs are missing");
assert!(path.exists(), "file must exist after write");
}
#[derive(Debug, serde::Serialize, serde::Deserialize, PartialEq)]
struct MockHydroModelSummary {
n_constant: usize,
n_fpha: usize,
total_planes: usize,
}
#[test]
fn write_and_read_hydro_model_summary_round_trips() {
let tmp = tempdir().expect("tempdir");
let path = tmp.path().join("training/hydro_models.json");
let summary = MockHydroModelSummary {
n_constant: 4,
n_fpha: 7,
total_planes: 35,
};
write_hydro_model_summary(&path, &summary).expect("write should succeed");
let decoded: MockHydroModelSummary =
read_hydro_model_summary(&path).expect("read should succeed");
assert_eq!(decoded, summary);
}
#[test]
fn hydro_model_summary_write_is_atomic_no_tmp_remains() {
let tmp = tempdir().expect("tempdir");
let path = tmp.path().join("hydro_models.json");
let summary = MockHydroModelSummary {
n_constant: 1,
n_fpha: 2,
total_planes: 6,
};
write_hydro_model_summary(&path, &summary).expect("write should succeed");
let tmp_path = path.with_extension("json.tmp");
assert!(
!tmp_path.exists(),
"tmp file should be removed after rename"
);
assert!(path.exists(), "final file should exist");
}
#[test]
fn read_hydro_model_summary_missing_file_is_not_found() {
let tmp = tempdir().expect("tempdir");
let path = tmp.path().join("does_not_exist.json");
let result = read_hydro_model_summary::<MockHydroModelSummary>(&path);
assert!(
matches!(
&result,
Err(OutputError::IoError { source, .. })
if source.kind() == std::io::ErrorKind::NotFound
),
"missing file must return IoError with NotFound kind so callers \
can degrade gracefully, got: {result:?}"
);
}
}