use std::path::Path;
use cobre_core::System;
use super::dictionary::write_dictionaries;
use super::error::OutputError;
use super::manifest::{
MetadataBounds, MetadataConfiguration, MetadataConvergence, MetadataIterations,
MetadataProblemDimensions, MetadataRowPool, MetadataScenarios, OutputContext,
SimulationMetadata, TrainingMetadata, write_simulation_metadata, write_training_metadata,
};
use super::parquet_config::ParquetWriterConfig;
use super::training_writer::TrainingParquetWriter;
use super::{SimulationOutput, TrainingOutput};
use crate::Config;
use crate::config::{ForwardPassesResolution, StoppingRuleConfig};
#[allow(clippy::cast_precision_loss, clippy::cast_possible_truncation)]
pub fn write_training_results(
output_dir: &Path,
training_output: &TrainingOutput,
system: &System,
config: &Config,
ctx: &OutputContext,
) -> Result<(), OutputError> {
std::fs::create_dir_all(output_dir.join("training/dictionaries"))
.map_err(|e| OutputError::io(output_dir.join("training/dictionaries"), e))?;
std::fs::create_dir_all(output_dir.join("training/timing"))
.map_err(|e| OutputError::io(output_dir.join("training/timing"), e))?;
std::fs::create_dir_all(output_dir.join("simulation"))
.map_err(|e| OutputError::io(output_dir.join("simulation"), e))?;
write_dictionaries(&output_dir.join("training/dictionaries"), system, config)?;
let parquet_config = ParquetWriterConfig::default();
let writer = TrainingParquetWriter::new(output_dir, &parquet_config)?;
writer.write(training_output)?;
let converged_at = training_output
.converged
.then_some(training_output.iterations_completed);
let max_iterations = extract_max_iterations(config);
let metadata = TrainingMetadata {
cobre_version: env!("CARGO_PKG_VERSION").to_string(),
hostname: ctx.hostname.clone(),
solver: ctx.solver.clone(),
solver_version: ctx.solver_version.clone(),
started_at: ctx.started_at.clone(),
completed_at: ctx.completed_at.clone(),
duration_seconds: training_output.total_time_ms as f64 / 1_000.0,
status: "complete".to_string(),
configuration: MetadataConfiguration {
seed: config.training.tree_seed,
max_iterations,
forward_passes: match config.resolve_forward_passes() {
Some(ForwardPassesResolution::Sampled(n)) => Some(n),
Some(ForwardPassesResolution::Enumerated) | None => None,
},
stopping_mode: config.training.stopping_mode.to_string(),
policy_mode: config.policy.mode.to_string(),
},
problem_dimensions: MetadataProblemDimensions {
num_stages: system.n_stages() as u32,
num_hydros: system.n_hydros() as u32,
num_thermals: system.n_thermals() as u32,
num_buses: system.n_buses() as u32,
num_lines: system.n_lines() as u32,
},
iterations: MetadataIterations {
completed: training_output.iterations_completed,
converged_at,
},
convergence: MetadataConvergence {
achieved: training_output.converged,
final_gap_percent: training_output.final_gap_percent,
termination_reason: training_output.termination_reason.clone(),
},
row_pool: MetadataRowPool {
total_generated: training_output.cut_stats.total_generated,
total_active: training_output.cut_stats.total_active,
peak_active: training_output.cut_stats.peak_active,
cuts_active: training_output.cut_stats.cuts_active,
rows_in_lp_total: training_output.cut_stats.rows_in_lp_total,
rows_in_lp_solve_count: training_output.cut_stats.rows_in_lp_solve_count,
rows_in_lp_max: training_output.cut_stats.rows_in_lp_max,
total_loaded: training_output.cut_stats.total_loaded,
},
bounds: MetadataBounds {
final_lower_bound: training_output.final_lower_bound,
final_upper_bound: training_output.final_upper_bound,
final_upper_bound_std: training_output.final_upper_bound_std,
final_upper_bound_kind: training_output.final_upper_bound_kind.clone(),
},
solve_stats: training_output.training_solve_stats.clone(),
setup: ctx.setup.clone(),
production_fit_deviation: ctx.production_fit_deviation.clone(),
distribution: ctx.distribution.clone(),
};
write_training_metadata(&output_dir.join("training/metadata.json"), &metadata)?;
std::fs::write(output_dir.join("training/_SUCCESS"), b"")
.map_err(|e| OutputError::io(output_dir.join("training/_SUCCESS"), e))?;
Ok(())
}
#[allow(clippy::cast_precision_loss)]
pub fn write_simulation_results(
output_dir: &Path,
simulation_output: &SimulationOutput,
ctx: &OutputContext,
) -> Result<(), OutputError> {
let metadata = SimulationMetadata {
cobre_version: env!("CARGO_PKG_VERSION").to_string(),
hostname: ctx.hostname.clone(),
solver: ctx.solver.clone(),
solver_version: ctx.solver_version.clone(),
started_at: ctx.started_at.clone(),
completed_at: ctx.completed_at.clone(),
duration_seconds: simulation_output.total_time_ms as f64 / 1_000.0,
status: "complete".to_string(),
scenarios: MetadataScenarios {
total: simulation_output.n_scenarios,
completed: simulation_output.completed,
failed: simulation_output.failed,
},
cost: simulation_output.cost.clone(),
solve_stats: simulation_output.solve_stats.clone(),
distribution: ctx.distribution.clone(),
};
write_simulation_metadata(&output_dir.join("simulation/metadata.json"), &metadata)?;
std::fs::write(output_dir.join("simulation/_SUCCESS"), b"")
.map_err(|e| OutputError::io(output_dir.join("simulation/_SUCCESS"), e))?;
Ok(())
}
pub fn write_results(
output_dir: &Path,
training_output: &TrainingOutput,
simulation_output: Option<&SimulationOutput>,
system: &System,
config: &Config,
ctx: &OutputContext,
) -> Result<(), OutputError> {
write_training_results(output_dir, training_output, system, config, ctx)?;
if let Some(sim) = simulation_output {
write_simulation_results(output_dir, sim, ctx)?;
}
Ok(())
}
fn extract_max_iterations(config: &Config) -> Option<u32> {
config
.training
.stopping_rules
.as_ref()?
.iter()
.find_map(|r| match r {
StoppingRuleConfig::IterationLimit { limit } => Some(*limit),
_ => None,
})
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::float_cmp,
clippy::cast_possible_truncation
)]
mod tests {
use super::*;
use crate::output::{IterationRecord, RowPoolStatistics, TrainingOutput};
use crate::{MetadataSimulationSolveStats, MetadataTrainingSolveStats};
use cobre_core::SystemBuilder;
fn make_iteration_record(iteration: u32) -> IterationRecord {
IterationRecord {
iteration,
lower_bound: 1.0,
upper_bound: 2.0,
upper_bound_std: 0.1,
gap_percent: Some(50.0),
cuts_added: 10,
cuts_removed: 2,
cuts_active: 8,
time_forward_ms: 100,
time_backward_ms: 200,
time_total_ms: 300,
forward_passes: 4,
lp_solves: 40,
time_forward_wall_ms: 100,
time_backward_wall_ms: 200,
time_cut_selection_ms: 0,
time_mpi_allreduce_ms: 0,
time_cut_sync_ms: 0,
time_lower_bound_ms: 0,
time_state_exchange_ms: 0,
time_cut_batch_build_ms: 0,
time_bwd_setup_ms: 0,
time_bwd_load_imbalance_ms: 0,
time_bwd_scheduling_overhead_ms: 0,
time_fwd_setup_ms: 0,
time_fwd_load_imbalance_ms: 0,
time_fwd_scheduling_overhead_ms: 0,
time_overhead_ms: 0,
solve_time_ms: 0.0,
mean_rows_in_lp: 0.0,
}
}
fn make_training_output(n_records: usize) -> TrainingOutput {
let records = (1..=n_records as u32).map(make_iteration_record).collect();
TrainingOutput {
convergence_records: records,
final_lower_bound: 99.5,
final_upper_bound: Some(101.0),
final_gap_percent: Some(1.51),
final_upper_bound_std: Some(0.5),
final_upper_bound_kind: "statistical".to_string(),
iterations_completed: n_records as u32,
converged: true,
termination_reason: "gap tolerance reached".to_string(),
total_time_ms: 5_000,
cut_stats: RowPoolStatistics {
total_generated: 200,
total_active: 80,
peak_active: 95,
cuts_active: 0,
rows_in_lp_total: 0,
rows_in_lp_solve_count: 0,
rows_in_lp_max: 0,
total_loaded: 0,
},
cut_selection_records: vec![],
worker_timing_records: vec![],
training_solve_stats: MetadataTrainingSolveStats::default(),
}
}
fn make_system() -> cobre_core::System {
SystemBuilder::new()
.build()
.expect("empty system must be valid")
}
fn make_config() -> crate::Config {
use crate::config::{
CheckpointingConfig, EstimationConfig, ExportsConfig, InflowNonNegativityConfig,
ModelingConfig, ParallelismConfig, PolicyConfig, PolicyMode, RowSelectionConfig,
SimulationConfig, StoppingMode, StoppingRuleConfig, TrainingConfig, TrainingSelection,
TrainingSolverConfig, UpperBoundEvaluationConfig,
};
crate::Config {
schema: None,
modeling: ModelingConfig {
inflow_non_negativity: InflowNonNegativityConfig::default(),
cost_scale_factor: None,
},
training: TrainingConfig {
enabled: true,
tree_seed: None,
stopping_rules: Some(vec![StoppingRuleConfig::IterationLimit { limit: 10 }]),
stopping_mode: StoppingMode::Any,
cut_selection: RowSelectionConfig::default(),
solver: TrainingSolverConfig::default(),
parallelism: ParallelismConfig::default(),
scenario_source: None,
selection: Some(TrainingSelection::Sampled { forward_passes: 4 }),
},
upper_bound_evaluation: UpperBoundEvaluationConfig::default(),
policy: PolicyConfig {
path: "./policy".to_string(),
mode: PolicyMode::Fresh,
checkpointing: CheckpointingConfig::default(),
boundary: None,
},
simulation: SimulationConfig {
enabled: false,
io_channel_capacity: 64,
scenario_source: None,
solver: None,
selection: None,
},
exports: ExportsConfig::default(),
estimation: EstimationConfig::default(),
}
}
fn make_simulation_output() -> SimulationOutput {
SimulationOutput {
n_scenarios: 10,
completed: 10,
failed: 0,
total_time_ms: 1_000,
partitions_written: vec!["simulation/costs/part-00.parquet".to_string()],
cost: None,
solve_stats: MetadataSimulationSolveStats::default(),
}
}
fn make_output_context() -> OutputContext {
use super::super::manifest::DistributionInfo;
OutputContext {
hostname: "test-host".to_string(),
solver: "highs".to_string(),
solver_version: None,
started_at: "2026-01-17T08:00:00Z".to_string(),
completed_at: "2026-01-17T12:30:00Z".to_string(),
distribution: DistributionInfo {
backend: "local".to_string(),
world_size: 1,
ranks_participated: 1,
num_hosts: 1,
threads_per_rank: 1,
mpi_library: None,
mpi_standard: None,
thread_level: None,
slurm_job_id: None,
hosts: Vec::new(),
},
setup: None,
production_fit_deviation: None,
}
}
#[test]
fn write_results_creates_training_directories() {
let tmp = tempfile::tempdir().unwrap();
let training = make_training_output(0);
write_results(
tmp.path(),
&training,
None,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_results must succeed");
assert!(tmp.path().join("training").is_dir(), "training/ must exist");
assert!(
tmp.path().join("training/dictionaries").is_dir(),
"training/dictionaries/ must exist"
);
assert!(
tmp.path().join("training/timing").is_dir(),
"training/timing/ must exist"
);
}
#[test]
fn write_results_creates_simulation_directory() {
let tmp = tempfile::tempdir().unwrap();
let training = make_training_output(0);
write_results(
tmp.path(),
&training,
None,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_results must succeed with simulation_output = None");
assert!(
tmp.path().join("simulation").is_dir(),
"simulation/ must exist even when simulation_output is None"
);
}
#[test]
fn write_results_returns_ok_on_success() {
let tmp = tempfile::tempdir().unwrap();
let training = make_training_output(3);
let simulation = SimulationOutput {
n_scenarios: 10,
completed: 10,
failed: 0,
total_time_ms: 1_500,
partitions_written: vec!["simulation/costs/part-00.parquet".to_string()],
cost: None,
solve_stats: MetadataSimulationSolveStats::default(),
};
let result = write_results(
tmp.path(),
&training,
Some(&simulation),
&make_system(),
&make_config(),
&make_output_context(),
);
assert!(
result.is_ok(),
"write_results must return Ok(()) on success"
);
}
#[test]
fn write_results_creates_success_marker() {
let tmp = tempfile::tempdir().unwrap();
let training = make_training_output(0);
write_results(
tmp.path(),
&training,
None,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_results must succeed");
assert!(
tmp.path().join("training/_SUCCESS").is_file(),
"training/_SUCCESS must exist after write_results"
);
}
#[test]
fn write_results_creates_metadata() {
let tmp = tempfile::tempdir().unwrap();
let training = make_training_output(0);
write_results(
tmp.path(),
&training,
None,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_results must succeed");
let path = tmp.path().join("training/metadata.json");
assert!(path.is_file(), "training/metadata.json must exist");
let content = std::fs::read_to_string(&path).unwrap();
let value: serde_json::Value =
serde_json::from_str(&content).expect("metadata.json must contain valid JSON");
assert_eq!(value["hostname"].as_str(), Some("test-host"));
assert_eq!(value["solver"].as_str(), Some("highs"));
assert!(value["started_at"].is_string());
assert!(value["completed_at"].is_string());
}
#[test]
fn write_results_metadata_row_pool_reports_loaded_boundary_cuts() {
let tmp = tempfile::tempdir().unwrap();
let mut training = make_training_output(0);
training.cut_stats.total_generated = 10_009;
training.cut_stats.total_active = 10_009;
training.cut_stats.total_loaded = 10_000;
write_results(
tmp.path(),
&training,
None,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_results must succeed");
let path = tmp.path().join("training/metadata.json");
let content = std::fs::read_to_string(&path).unwrap();
let value: serde_json::Value =
serde_json::from_str(&content).expect("metadata.json must contain valid JSON");
assert_eq!(value["row_pool"]["total_generated"].as_u64(), Some(10_009));
assert_eq!(value["row_pool"]["total_loaded"].as_u64(), Some(10_000));
}
#[test]
fn write_results_creates_convergence_parquet() {
let tmp = tempfile::tempdir().unwrap();
let training = make_training_output(3);
write_results(
tmp.path(),
&training,
None,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_results must succeed");
assert!(
tmp.path().join("training/convergence.parquet").is_file(),
"training/convergence.parquet must exist"
);
}
#[test]
fn write_results_convergence_parquet_row_count() {
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
let tmp = tempfile::tempdir().unwrap();
let training = make_training_output(3);
write_results(
tmp.path(),
&training,
None,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_results must succeed");
let path = tmp.path().join("training/convergence.parquet");
let file = std::fs::File::open(&path).unwrap();
let reader = ParquetRecordBatchReaderBuilder::try_new(file)
.unwrap()
.build()
.unwrap();
let total_rows: usize = reader
.map(|b| b.expect("batch must be Ok").num_rows())
.sum();
assert_eq!(total_rows, 3, "convergence.parquet must have 3 rows");
}
#[test]
fn write_results_empty_training_convergence_parquet_correct_schema() {
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
let tmp = tempfile::tempdir().unwrap();
let training = make_training_output(0);
write_results(
tmp.path(),
&training,
None,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_results must succeed");
let path = tmp.path().join("training/convergence.parquet");
assert!(path.is_file(), "training/convergence.parquet must exist");
let file = std::fs::File::open(&path).unwrap();
let builder = ParquetRecordBatchReaderBuilder::try_new(file).unwrap();
let schema = builder.schema().clone();
let reader = builder.build().unwrap();
let total_rows: usize = reader
.map(|b| b.expect("batch must be Ok").num_rows())
.sum();
assert_eq!(total_rows, 0, "empty training must produce 0 rows");
assert_eq!(
schema.fields().len(),
15,
"convergence schema must have 15 columns"
);
assert!(
tmp.path().join("training/_SUCCESS").is_file(),
"training/_SUCCESS must exist even with 0 records"
);
}
#[test]
fn write_results_simulation_success_marker_conditional() {
let tmp = tempfile::tempdir().unwrap();
let training = make_training_output(0);
let simulation = SimulationOutput {
n_scenarios: 10,
completed: 10,
failed: 0,
total_time_ms: 0,
partitions_written: vec![],
cost: None,
solve_stats: MetadataSimulationSolveStats::default(),
};
write_results(
tmp.path(),
&training,
Some(&simulation),
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_results must succeed");
assert!(
tmp.path().join("simulation/_SUCCESS").is_file(),
"simulation/_SUCCESS must exist when simulation_output is Some"
);
assert!(
tmp.path().join("training/_SUCCESS").is_file(),
"training/_SUCCESS must exist"
);
let tmp2 = tempfile::tempdir().unwrap();
write_results(
tmp2.path(),
&training,
None,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_results must succeed");
assert!(
!tmp2.path().join("simulation/_SUCCESS").exists(),
"simulation/_SUCCESS must NOT exist when simulation_output is None"
);
}
#[test]
fn write_results_simulation_metadata_scenarios_total() {
let tmp = tempfile::tempdir().unwrap();
let training = make_training_output(3);
let simulation = SimulationOutput {
n_scenarios: 10,
completed: 10,
failed: 0,
total_time_ms: 0,
partitions_written: vec![],
cost: None,
solve_stats: MetadataSimulationSolveStats::default(),
};
write_results(
tmp.path(),
&training,
Some(&simulation),
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_results must succeed");
let path = tmp.path().join("simulation/metadata.json");
assert!(path.is_file(), "simulation/metadata.json must exist");
let content = std::fs::read_to_string(&path).unwrap();
let value: serde_json::Value = serde_json::from_str(&content).unwrap();
assert_eq!(
value["scenarios"]["total"].as_u64(),
Some(10),
"$.scenarios.total must equal 10"
);
}
#[test]
fn write_results_creates_dictionaries() {
let tmp = tempfile::tempdir().unwrap();
let training = make_training_output(0);
write_results(
tmp.path(),
&training,
None,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_results must succeed");
assert!(
tmp.path()
.join("training/dictionaries/codes.json")
.is_file(),
"training/dictionaries/codes.json must exist"
);
}
#[test]
fn write_results_codes_json_contains_operative_state() {
let tmp = tempfile::tempdir().unwrap();
let training = make_training_output(0);
write_results(
tmp.path(),
&training,
None,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_results must succeed");
let path = tmp.path().join("training/dictionaries/codes.json");
let content = std::fs::read_to_string(&path).unwrap();
let value: serde_json::Value = serde_json::from_str(&content).unwrap();
assert!(
value["operative_state"].is_object(),
"codes.json must contain an operative_state object"
);
assert_eq!(
value["operative_state"]["2"].as_str(),
Some("operating"),
r#"codes.json operative_state["2"] must equal "operating""#
);
}
#[test]
fn write_training_results_produces_complete_output() {
let tmp = tempfile::tempdir().unwrap();
let training = make_training_output(3);
write_training_results(
tmp.path(),
&training,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_training_results must succeed");
assert!(tmp.path().join("training").is_dir());
assert!(tmp.path().join("training/dictionaries").is_dir());
assert!(tmp.path().join("training/timing").is_dir());
assert!(tmp.path().join("training/metadata.json").is_file());
assert!(tmp.path().join("training/_SUCCESS").is_file());
assert!(
tmp.path().join("simulation").is_dir(),
"simulation/ directory must be created by write_training_results"
);
}
#[test]
fn write_simulation_results_produces_metadata_and_success() {
let tmp = tempfile::tempdir().unwrap();
std::fs::create_dir_all(tmp.path().join("simulation")).unwrap();
let sim = make_simulation_output();
write_simulation_results(tmp.path(), &sim, &make_output_context())
.expect("write_simulation_results must succeed");
assert!(tmp.path().join("simulation/metadata.json").is_file());
assert!(tmp.path().join("simulation/_SUCCESS").is_file());
}
#[test]
fn split_functions_match_write_results_output() {
let tmp_combined = tempfile::tempdir().unwrap();
let tmp_split = tempfile::tempdir().unwrap();
let training = make_training_output(2);
let sim = make_simulation_output();
let ctx = make_output_context();
write_results(
tmp_combined.path(),
&training,
Some(&sim),
&make_system(),
&make_config(),
&ctx,
)
.expect("write_results must succeed");
write_training_results(
tmp_split.path(),
&training,
&make_system(),
&make_config(),
&ctx,
)
.expect("write_training_results must succeed");
write_simulation_results(tmp_split.path(), &sim, &ctx)
.expect("write_simulation_results must succeed");
let combined_training_success = tmp_combined.path().join("training/_SUCCESS").is_file();
let split_training_success = tmp_split.path().join("training/_SUCCESS").is_file();
assert_eq!(combined_training_success, split_training_success);
let combined_sim_success = tmp_combined.path().join("simulation/_SUCCESS").is_file();
let split_sim_success = tmp_split.path().join("simulation/_SUCCESS").is_file();
assert_eq!(combined_sim_success, split_sim_success);
let combined_metadata = tmp_combined.path().join("training/metadata.json").is_file();
let split_metadata = tmp_split.path().join("training/metadata.json").is_file();
assert_eq!(combined_metadata, split_metadata);
let combined_sim_metadata = tmp_combined
.path()
.join("simulation/metadata.json")
.is_file();
let split_sim_metadata = tmp_split.path().join("simulation/metadata.json").is_file();
assert_eq!(combined_sim_metadata, split_sim_metadata);
}
#[test]
fn extract_max_iterations_from_config() {
let config = make_config();
assert_eq!(extract_max_iterations(&config), Some(10));
}
#[test]
fn training_metadata_has_max_iterations() {
let tmp = tempfile::tempdir().unwrap();
let training = make_training_output(0);
write_training_results(
tmp.path(),
&training,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_training_results must succeed");
let path = tmp.path().join("training/metadata.json");
let content = std::fs::read_to_string(&path).unwrap();
let value: serde_json::Value = serde_json::from_str(&content).unwrap();
assert_eq!(
value["configuration"]["max_iterations"].as_u64(),
Some(10),
"configuration.max_iterations must be extracted from stopping rules"
);
}
fn read_convergence_kind_and_std_nulls(dir: &std::path::Path) -> (String, usize) {
use arrow::array::{Array, Float64Array, StringArray};
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
let path = dir.join("training/convergence.parquet");
let file = std::fs::File::open(&path).unwrap();
let mut reader = ParquetRecordBatchReaderBuilder::try_new(file)
.unwrap()
.build()
.unwrap();
let batch = reader.next().unwrap().unwrap();
let kind = batch
.column_by_name("upper_bound_kind")
.unwrap()
.as_any()
.downcast_ref::<StringArray>()
.unwrap()
.value(0)
.to_string();
let std_nulls = batch
.column_by_name("upper_bound_std")
.unwrap()
.as_any()
.downcast_ref::<Float64Array>()
.unwrap()
.null_count();
(kind, std_nulls)
}
#[test]
fn exact_bound_writes_exact_kind_and_null_std_in_both_artifacts() {
use crate::output::manifest::read_training_metadata;
let tmp = tempfile::tempdir().unwrap();
let mut training = make_training_output(3);
training.final_upper_bound_kind = "exact".to_string();
training.final_upper_bound_std = None;
write_training_results(
tmp.path(),
&training,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write must succeed");
let (kind, std_nulls) = read_convergence_kind_and_std_nulls(tmp.path());
assert_eq!(kind, "exact", "convergence upper_bound_kind must be exact");
assert_eq!(
std_nulls, 3,
"every upper_bound_std must be NULL under exact"
);
let metadata = read_training_metadata(&tmp.path().join("training/metadata.json")).unwrap();
assert_eq!(metadata.bounds.final_upper_bound_kind, "exact");
assert_eq!(metadata.bounds.final_upper_bound_std, None);
}
#[test]
fn statistical_bound_writes_statistical_kind_and_populated_std_in_both_artifacts() {
use crate::output::manifest::read_training_metadata;
let tmp = tempfile::tempdir().unwrap();
let mut training = make_training_output(3);
training.final_upper_bound_kind = "statistical".to_string();
training.final_upper_bound_std = Some(0.5);
write_training_results(
tmp.path(),
&training,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write must succeed");
let (kind, std_nulls) = read_convergence_kind_and_std_nulls(tmp.path());
assert_eq!(kind, "statistical");
assert_eq!(
std_nulls, 0,
"upper_bound_std must be populated under statistical"
);
let metadata = read_training_metadata(&tmp.path().join("training/metadata.json")).unwrap();
assert_eq!(metadata.bounds.final_upper_bound_kind, "statistical");
assert_eq!(metadata.bounds.final_upper_bound_std, Some(0.5));
}
#[test]
fn training_results_persist_bounds_and_solve_stats() {
use crate::output::manifest::read_training_metadata;
let tmp = tempfile::tempdir().unwrap();
let mut training = make_training_output(2);
training.final_lower_bound = 48_500.0;
training.final_upper_bound = Some(49_000.0);
training.final_upper_bound_std = Some(250.0);
training.training_solve_stats = MetadataTrainingSolveStats {
total_lp_solves: Some(120),
first_try: Some(110),
retried: Some(10),
failed: Some(0),
forward_solve_seconds: Some(3.0),
backward_solve_seconds: Some(5.0),
parallelism: Some(4),
};
write_training_results(
tmp.path(),
&training,
&make_system(),
&make_config(),
&make_output_context(),
)
.expect("write_training_results must succeed");
let metadata = read_training_metadata(&tmp.path().join("training/metadata.json"))
.expect("read_training_metadata must succeed");
assert_eq!(metadata.bounds.final_lower_bound, 48_500.0);
assert_eq!(metadata.bounds.final_upper_bound, Some(49_000.0));
assert_eq!(metadata.bounds.final_upper_bound_std, Some(250.0));
assert_eq!(metadata.solve_stats.total_lp_solves, Some(120));
assert_eq!(metadata.solve_stats.forward_solve_seconds, Some(3.0));
assert_eq!(metadata.solve_stats.backward_solve_seconds, Some(5.0));
assert_eq!(metadata.solve_stats.parallelism, Some(4));
}
#[test]
fn simulation_results_persist_cost_and_solve_stats() {
use crate::output::manifest::read_simulation_metadata;
let tmp = tempfile::tempdir().unwrap();
std::fs::create_dir_all(tmp.path().join("simulation")).unwrap();
let mut sim = make_simulation_output();
sim.cost = Some(crate::MetadataCost {
mean_cost: 12_345.6,
std_cost: 678.9,
});
sim.solve_stats = MetadataSimulationSolveStats {
total_lp_solves: Some(200),
first_try: Some(190),
retried: Some(9),
failed: Some(1),
solve_seconds: Some(7.5),
parallelism: Some(8),
};
write_simulation_results(tmp.path(), &sim, &make_output_context())
.expect("write_simulation_results must succeed");
let metadata = read_simulation_metadata(&tmp.path().join("simulation/metadata.json"))
.expect("read_simulation_metadata must succeed");
let cost = metadata.cost.expect("cost must be persisted");
assert_eq!(cost.mean_cost, 12_345.6);
assert_eq!(cost.std_cost, 678.9);
assert_eq!(metadata.solve_stats.total_lp_solves, Some(200));
assert_eq!(metadata.solve_stats.solve_seconds, Some(7.5));
assert_eq!(metadata.solve_stats.parallelism, Some(8));
}
}