use crate::LoadError;
use serde::{Deserialize, Serialize};
use std::path::Path;
#[derive(Debug, Clone, Deserialize, Serialize)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct Config {
#[serde(rename = "$schema")]
pub schema: Option<String>,
#[serde(default)]
pub modeling: ModelingConfig,
pub training: TrainingConfig,
#[serde(default)]
pub upper_bound_evaluation: UpperBoundEvaluationConfig,
#[serde(default)]
pub policy: PolicyConfig,
#[serde(default)]
pub simulation: SimulationConfig,
#[serde(default)]
pub exports: ExportsConfig,
#[serde(default)]
pub estimation: EstimationConfig,
}
#[derive(Debug, Clone, Deserialize, Serialize, Default)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct ModelingConfig {
#[serde(default)]
pub inflow_non_negativity: InflowNonNegativityConfig,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(default)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct InflowNonNegativityConfig {
pub method: String,
pub penalty_cost: f64,
}
impl Default for InflowNonNegativityConfig {
fn default() -> Self {
Self {
method: "penalty".to_string(),
penalty_cost: 1000.0,
}
}
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct TrainingConfig {
#[serde(default = "TrainingConfig::default_enabled")]
pub enabled: bool,
#[serde(default)]
pub seed: Option<i64>,
pub forward_passes: Option<u32>,
pub stopping_rules: Option<Vec<StoppingRuleConfig>>,
#[serde(default = "TrainingConfig::default_stopping_mode")]
pub stopping_mode: String,
#[serde(default)]
pub cut_formulation: Option<String>,
#[serde(default)]
pub forward_pass: Option<ForwardPassConfig>,
#[serde(default)]
pub cut_selection: CutSelectionConfig,
#[serde(default)]
pub solver: TrainingSolverConfig,
}
impl TrainingConfig {
fn default_enabled() -> bool {
true
}
fn default_stopping_mode() -> String {
"any".to_string()
}
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct ForwardPassConfig {
#[serde(rename = "type")]
pub pass_type: String,
}
#[derive(Debug, Clone, Deserialize, Serialize, Default)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct CutSelectionConfig {
#[serde(default)]
pub enabled: Option<bool>,
#[serde(default)]
pub method: Option<String>,
#[serde(default)]
pub threshold: Option<u32>,
#[serde(default)]
pub check_frequency: Option<u32>,
#[serde(default)]
pub cut_activity_tolerance: Option<f64>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(default)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct TrainingSolverConfig {
pub retry_max_attempts: u32,
pub retry_time_budget_seconds: f64,
}
impl Default for TrainingSolverConfig {
fn default() -> Self {
Self {
retry_max_attempts: 5,
retry_time_budget_seconds: 30.0,
}
}
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub enum StoppingRuleConfig {
IterationLimit {
limit: u32,
},
TimeLimit {
seconds: f64,
},
BoundStalling {
iterations: u32,
tolerance: f64,
},
Simulation {
replications: u32,
period: u32,
bound_window: u32,
distance_tol: f64,
bound_tol: f64,
},
}
#[derive(Debug, Clone, Deserialize, Serialize, Default)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct UpperBoundEvaluationConfig {
#[serde(default)]
pub enabled: Option<bool>,
#[serde(default)]
pub initial_iteration: Option<u32>,
#[serde(default)]
pub interval_iterations: Option<u32>,
#[serde(default)]
pub lipschitz: LipschitzConfig,
}
#[derive(Debug, Clone, Deserialize, Serialize, Default)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct LipschitzConfig {
#[serde(default)]
pub mode: Option<String>,
#[serde(default)]
pub fallback_value: Option<f64>,
#[serde(default)]
pub scale_factor: Option<f64>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(default)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct PolicyConfig {
pub path: String,
pub mode: String,
pub validate_compatibility: bool,
pub checkpointing: CheckpointingConfig,
}
impl Default for PolicyConfig {
fn default() -> Self {
Self {
path: "./policy".to_string(),
mode: "fresh".to_string(),
validate_compatibility: true,
checkpointing: CheckpointingConfig::default(),
}
}
}
#[derive(Debug, Clone, Deserialize, Serialize, Default)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct CheckpointingConfig {
#[serde(default)]
pub enabled: Option<bool>,
#[serde(default)]
pub initial_iteration: Option<u32>,
#[serde(default)]
pub interval_iterations: Option<u32>,
#[serde(default)]
pub store_basis: Option<bool>,
#[serde(default)]
pub compress: Option<bool>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(default)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct SimulationConfig {
pub enabled: bool,
pub num_scenarios: u32,
pub policy_type: String,
pub output_path: Option<String>,
pub output_mode: Option<String>,
pub io_channel_capacity: u32,
pub sampling_scheme: SimulationSamplingConfig,
}
impl Default for SimulationConfig {
fn default() -> Self {
Self {
enabled: false,
num_scenarios: 2000,
policy_type: "outer".to_string(),
output_path: None,
output_mode: None,
io_channel_capacity: 64,
sampling_scheme: SimulationSamplingConfig::default(),
}
}
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(default)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct SimulationSamplingConfig {
#[serde(rename = "type")]
pub scheme_type: String,
}
impl Default for SimulationSamplingConfig {
fn default() -> Self {
Self {
scheme_type: "in_sample".to_string(),
}
}
}
#[derive(Debug, Clone, Deserialize, Serialize, Default)]
#[serde(rename_all = "snake_case")]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub enum OrderSelectionMethod {
Fixed,
#[default]
Aic,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(default)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct EstimationConfig {
pub max_order: u32,
pub order_selection: OrderSelectionMethod,
pub min_observations_per_season: u32,
}
impl Default for EstimationConfig {
fn default() -> Self {
Self {
max_order: 6,
order_selection: OrderSelectionMethod::Aic,
min_observations_per_season: 30,
}
}
}
#[allow(clippy::struct_excessive_bools)]
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(default)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct ExportsConfig {
pub training: bool,
pub cuts: bool,
pub states: bool,
pub vertices: bool,
pub simulation: bool,
pub forward_detail: bool,
pub backward_detail: bool,
pub compression: Option<String>,
}
impl Default for ExportsConfig {
fn default() -> Self {
Self {
training: true,
cuts: true,
states: true,
vertices: true,
simulation: true,
forward_detail: false,
backward_detail: false,
compression: None,
}
}
}
pub fn parse_config(path: &Path) -> Result<Config, LoadError> {
let raw = std::fs::read_to_string(path).map_err(|e| LoadError::io(path, e))?;
let config: Config = serde_json::from_str(&raw).map_err(|e| {
let msg = e.to_string();
if msg.contains("unknown variant") || msg.contains("missing field") {
LoadError::SchemaError {
path: path.to_path_buf(),
field: extract_field_from_serde_msg(&msg),
message: msg,
}
} else {
LoadError::parse(path, msg)
}
})?;
validate_config(&config, path)?;
Ok(config)
}
fn extract_field_from_serde_msg(msg: &str) -> String {
if let Some(start) = msg.find('`') {
if let Some(end) = msg[start + 1..].find('`') {
return msg[start + 1..start + 1 + end].to_string();
}
}
"<unknown>".to_string()
}
fn validate_config(config: &Config, path: &Path) -> Result<(), LoadError> {
if config.training.forward_passes.is_none() {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: "training.forward_passes".to_string(),
message: "required field is missing".to_string(),
});
}
if config.training.stopping_rules.is_none() {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: "training.stopping_rules".to_string(),
message: "required field is missing".to_string(),
});
}
Ok(())
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::panic,
clippy::too_many_lines,
clippy::doc_markdown
)]
mod tests {
use super::*;
use std::io::Write;
use tempfile::NamedTempFile;
fn write_config(content: &str) -> NamedTempFile {
let mut f = NamedTempFile::new().unwrap();
f.write_all(content.as_bytes()).unwrap();
f
}
#[test]
fn test_parse_minimal_config() {
let f = write_config(
r#"{"training": {"seed": 42, "forward_passes": 192, "stopping_rules": [{"type": "iteration_limit", "limit": 50}]}}"#,
);
let cfg = parse_config(f.path()).unwrap();
assert_eq!(cfg.training.forward_passes, Some(192));
assert_eq!(cfg.training.seed, Some(42));
assert_eq!(cfg.training.stopping_mode, "any");
assert!(cfg.training.enabled);
assert_eq!(
cfg.modeling.inflow_non_negativity.method,
"penalty".to_string()
);
assert!((cfg.modeling.inflow_non_negativity.penalty_cost - 1000.0).abs() < f64::EPSILON);
assert!(!cfg.simulation.enabled);
assert_eq!(cfg.simulation.num_scenarios, 2000);
assert_eq!(cfg.policy.mode, "fresh");
assert_eq!(cfg.policy.path, "./policy");
assert!(cfg.policy.validate_compatibility);
assert!(cfg.exports.training);
assert!(cfg.exports.cuts);
}
#[test]
fn test_missing_forward_passes() {
let f = write_config(
r#"{"training": {"seed": 1, "stopping_rules": [{"type": "iteration_limit", "limit": 10}]}}"#,
);
let err = parse_config(f.path()).unwrap_err();
match &err {
LoadError::SchemaError { field, .. } => {
assert!(
field.contains("forward_passes"),
"field should contain 'forward_passes', got: {field}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn test_missing_stopping_rules() {
let f = write_config(r#"{"training": {"seed": 1, "forward_passes": 100}}"#);
let err = parse_config(f.path()).unwrap_err();
match &err {
LoadError::SchemaError { field, .. } => {
assert!(
field.contains("stopping_rules"),
"field should contain 'stopping_rules', got: {field}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn test_nonexistent_file() {
let path = std::path::Path::new("/nonexistent/path/config.json");
let err = parse_config(path).unwrap_err();
match &err {
LoadError::IoError { path: p, .. } => {
assert_eq!(p, path);
}
other => panic!("expected IoError, got: {other:?}"),
}
}
#[test]
fn test_parse_full_config() {
let json = r#"{
"$schema": "https://raw.githubusercontent.com/cobre-rs/cobre/refs/heads/main/book/src/schemas/config.schema.json",
"modeling": {
"inflow_non_negativity": {
"method": "penalty",
"penalty_cost": 500.0
}
},
"training": {
"seed": 42,
"forward_passes": 192,
"stopping_rules": [
{"type": "iteration_limit", "limit": 50},
{"type": "bound_stalling", "iterations": 10, "tolerance": 0.0001}
],
"stopping_mode": "any",
"cut_formulation": "single",
"forward_pass": {"type": "default"},
"cut_selection": {
"enabled": true,
"method": "domination",
"threshold": 0
}
},
"upper_bound_evaluation": {
"enabled": true,
"initial_iteration": 10,
"interval_iterations": 5
},
"policy": {
"path": "./policy",
"mode": "fresh",
"checkpointing": {
"enabled": true,
"initial_iteration": 10,
"interval_iterations": 10,
"store_basis": true,
"compress": true
},
"validate_compatibility": true
},
"simulation": {
"enabled": true,
"num_scenarios": 2000,
"policy_type": "outer",
"output_path": "./simulation",
"output_mode": "streaming",
"sampling_scheme": {"type": "in_sample"}
},
"exports": {
"training": true,
"cuts": true,
"states": true,
"vertices": true,
"simulation": true,
"forward_detail": false,
"backward_detail": false,
"compression": "zstd"
}
}"#;
let f = write_config(json);
let cfg = parse_config(f.path()).unwrap();
assert_eq!(cfg.modeling.inflow_non_negativity.method, "penalty");
assert!((cfg.modeling.inflow_non_negativity.penalty_cost - 500.0).abs() < f64::EPSILON);
assert_eq!(cfg.training.forward_passes, Some(192));
assert_eq!(cfg.training.stopping_mode, "any");
let rules = cfg.training.stopping_rules.as_ref().unwrap();
assert_eq!(rules.len(), 2);
assert_eq!(cfg.training.cut_formulation.as_deref(), Some("single"));
let cut_sel = &cfg.training.cut_selection;
assert_eq!(cut_sel.enabled, Some(true));
assert_eq!(cut_sel.method.as_deref(), Some("domination"));
assert_eq!(cfg.upper_bound_evaluation.enabled, Some(true));
assert_eq!(cfg.upper_bound_evaluation.initial_iteration, Some(10));
assert_eq!(cfg.policy.mode, "fresh");
assert!(cfg.policy.validate_compatibility);
assert_eq!(cfg.policy.checkpointing.enabled, Some(true));
assert!(cfg.simulation.enabled);
assert_eq!(cfg.simulation.num_scenarios, 2000);
assert_eq!(cfg.simulation.policy_type, "outer");
assert!(cfg.exports.training);
assert_eq!(cfg.exports.compression.as_deref(), Some("zstd"));
assert!(!cfg.exports.forward_detail);
}
#[test]
fn test_invalid_json_syntax() {
let f = write_config(r#"{"training": {not valid json}}"#);
let err = parse_config(f.path()).unwrap_err();
assert!(
matches!(err, LoadError::ParseError { .. }),
"expected ParseError, got: {err:?}"
);
}
#[test]
fn test_stopping_rule_variants() {
let json = r#"{
"training": {
"forward_passes": 10,
"stopping_rules": [
{"type": "iteration_limit", "limit": 100},
{"type": "time_limit", "seconds": 3600.0},
{"type": "bound_stalling", "iterations": 10, "tolerance": 0.0001},
{
"type": "simulation",
"replications": 100,
"period": 20,
"bound_window": 5,
"distance_tol": 0.01,
"bound_tol": 0.0001
}
]
}
}"#;
let f = write_config(json);
let cfg = parse_config(f.path()).unwrap();
let rules = cfg.training.stopping_rules.unwrap();
assert_eq!(rules.len(), 4);
assert!(matches!(
rules[0],
StoppingRuleConfig::IterationLimit { limit: 100 }
));
assert!(
matches!(rules[1], StoppingRuleConfig::TimeLimit { seconds } if (seconds - 3600.0).abs() < f64::EPSILON)
);
assert!(matches!(
rules[2],
StoppingRuleConfig::BoundStalling { iterations: 10, .. }
));
assert!(matches!(
rules[3],
StoppingRuleConfig::Simulation {
replications: 100,
period: 20,
..
}
));
}
#[test]
fn test_unknown_stopping_rule_type() {
let f = write_config(
r#"{"training": {"forward_passes": 10, "stopping_rules": [{"type": "nonexistent_rule"}]}}"#,
);
let err = parse_config(f.path()).unwrap_err();
assert!(
matches!(err, LoadError::SchemaError { .. }),
"expected SchemaError for unknown rule type, got: {err:?}"
);
}
#[test]
fn test_config_has_no_version_field() {
let f = write_config(
r#"{"training": {"forward_passes": 1, "stopping_rules": [{"type": "iteration_limit", "limit": 10}]}}"#,
);
let cfg = parse_config(f.path()).unwrap();
assert!(cfg.schema.is_none(), "schema should be None when absent");
}
#[test]
fn test_schema_field_accepted() {
let f = write_config(
r#"{
"$schema": "https://raw.githubusercontent.com/cobre-rs/cobre/refs/heads/main/book/src/schemas/config.schema.json",
"training": {
"forward_passes": 1,
"stopping_rules": [{"type": "iteration_limit", "limit": 10}]
}
}"#,
);
let cfg = parse_config(f.path()).unwrap();
assert_eq!(
cfg.schema.as_deref(),
Some(
"https://raw.githubusercontent.com/cobre-rs/cobre/refs/heads/main/book/src/schemas/config.schema.json"
),
"schema field should be stored when present in JSON"
);
}
#[test]
fn test_legacy_version_field_silently_ignored() {
let f = write_config(
r#"{
"version": "1.0.0",
"training": {
"forward_passes": 1,
"stopping_rules": [{"type": "iteration_limit", "limit": 10}]
}
}"#,
);
let cfg = parse_config(f.path()).unwrap();
assert_eq!(cfg.training.forward_passes, Some(1));
}
#[test]
fn test_truncation_method_accepted() {
let f = write_config(
r#"{
"modeling": {
"inflow_non_negativity": {
"method": "truncation"
}
},
"training": {
"forward_passes": 10,
"stopping_rules": [{"type": "iteration_limit", "limit": 5}]
}
}"#,
);
let cfg = parse_config(f.path()).unwrap();
assert_eq!(
cfg.modeling.inflow_non_negativity.method, "truncation",
"method field should round-trip as 'truncation'"
);
assert!(
(cfg.modeling.inflow_non_negativity.penalty_cost - 1000.0).abs() < f64::EPSILON,
"penalty_cost should be the default 1000.0 when absent from JSON"
);
}
#[test]
fn test_estimation_config_defaults() {
let f = write_config(
r#"{"training": {"forward_passes": 10, "stopping_rules": [{"type": "iteration_limit", "limit": 5}]}}"#,
);
let cfg = parse_config(f.path()).unwrap();
assert_eq!(cfg.estimation.max_order, 6);
assert!(
matches!(cfg.estimation.order_selection, OrderSelectionMethod::Aic),
"default order_selection should be Aic"
);
assert_eq!(cfg.estimation.min_observations_per_season, 30);
}
#[test]
fn test_estimation_config_explicit() {
let f = write_config(
r#"{
"training": {"forward_passes": 10, "stopping_rules": [{"type": "iteration_limit", "limit": 5}]},
"estimation": {"max_order": 3, "order_selection": "fixed", "min_observations_per_season": 20}
}"#,
);
let cfg = parse_config(f.path()).unwrap();
assert_eq!(cfg.estimation.max_order, 3);
assert!(
matches!(cfg.estimation.order_selection, OrderSelectionMethod::Fixed),
"order_selection should be Fixed"
);
assert_eq!(cfg.estimation.min_observations_per_season, 20);
}
#[test]
fn test_estimation_config_unknown_order_selection() {
let f = write_config(
r#"{
"training": {"forward_passes": 10, "stopping_rules": [{"type": "iteration_limit", "limit": 5}]},
"estimation": {"order_selection": "bogus"}
}"#,
);
let err = parse_config(f.path()).unwrap_err();
match &err {
LoadError::SchemaError { message, .. } => {
assert!(
message.contains("unknown variant"),
"message should contain 'unknown variant', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
}