use cobre_core::{
EntityId,
entities::{ContractType, EnergyContract},
};
use serde::Deserialize;
use std::collections::HashSet;
use std::path::Path;
use super::parse_operational_start_date;
use crate::LoadError;
#[derive(Deserialize)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
#[serde(deny_unknown_fields)]
pub(crate) struct RawContractFile {
#[serde(rename = "$schema")]
_schema: Option<String>,
contracts: Vec<RawContract>,
}
#[derive(Deserialize)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
#[serde(deny_unknown_fields)]
pub(crate) struct RawContract {
id: i32,
name: String,
operational_start_date: String,
bus_id: i32,
#[serde(rename = "type")]
contract_type: RawContractType,
#[serde(default)]
entry_stage_id: Option<i32>,
#[serde(default)]
exit_stage_id: Option<i32>,
price_per_mwh: f64,
limits: RawContractLimits,
}
#[derive(Deserialize)]
#[serde(rename_all = "snake_case")]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub(crate) enum RawContractType {
Import,
Export,
}
#[derive(Deserialize)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
#[serde(deny_unknown_fields)]
pub(crate) struct RawContractLimits {
min_mw: f64,
max_mw: f64,
}
pub fn parse_energy_contracts(path: &Path) -> Result<Vec<EnergyContract>, LoadError> {
let raw_text = std::fs::read_to_string(path).map_err(|e| LoadError::io(path, e))?;
let raw: RawContractFile =
serde_json::from_str(&raw_text).map_err(|e| LoadError::parse(path, e.to_string()))?;
validate_raw_contracts(&raw, path)?;
convert_contracts(raw, path)
}
fn validate_raw_contracts(raw: &RawContractFile, path: &Path) -> Result<(), LoadError> {
validate_no_duplicate_contract_ids(&raw.contracts, path)?;
for (i, contract) in raw.contracts.iter().enumerate() {
validate_contract_limits(&contract.limits, i, path)?;
}
Ok(())
}
fn validate_no_duplicate_contract_ids(
contracts: &[RawContract],
path: &Path,
) -> Result<(), LoadError> {
let mut seen: HashSet<i32> = HashSet::new();
for (i, contract) in contracts.iter().enumerate() {
if !seen.insert(contract.id) {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("contracts[{i}].id"),
message: format!("duplicate id {} in contracts array", contract.id),
});
}
}
Ok(())
}
fn validate_contract_limits(
limits: &RawContractLimits,
contract_index: usize,
path: &Path,
) -> Result<(), LoadError> {
if limits.min_mw < 0.0 {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("contracts[{contract_index}].limits.min_mw"),
message: format!("limits.min_mw must be >= 0.0, got {}", limits.min_mw),
});
}
if limits.max_mw < 0.0 {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("contracts[{contract_index}].limits.max_mw"),
message: format!("limits.max_mw must be >= 0.0, got {}", limits.max_mw),
});
}
if limits.max_mw < limits.min_mw {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("contracts[{contract_index}].limits.max_mw"),
message: format!(
"limits.max_mw ({}) must be >= limits.min_mw ({})",
limits.max_mw, limits.min_mw
),
});
}
Ok(())
}
fn convert_contracts(raw: RawContractFile, path: &Path) -> Result<Vec<EnergyContract>, LoadError> {
let mut contracts: Vec<EnergyContract> = raw
.contracts
.into_iter()
.enumerate()
.map(|(i, raw_contract)| {
let operational_start_date = parse_operational_start_date(
&raw_contract.operational_start_date,
path,
&format!("contracts[{i}].operational_start_date"),
)?;
let contract_type = match raw_contract.contract_type {
RawContractType::Import => ContractType::Import,
RawContractType::Export => ContractType::Export,
};
Ok(EnergyContract {
id: EntityId(raw_contract.id),
name: raw_contract.name,
operational_start_date,
bus_id: EntityId(raw_contract.bus_id),
contract_type,
entry_stage_id: raw_contract.entry_stage_id,
exit_stage_id: raw_contract.exit_stage_id,
price_per_mwh: raw_contract.price_per_mwh,
min_mw: raw_contract.limits.min_mw,
max_mw: raw_contract.limits.max_mw,
})
})
.collect::<Result<_, LoadError>>()?;
contracts.sort_by_key(|c| c.id.0);
Ok(contracts)
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::panic, clippy::too_many_lines)]
mod tests {
use super::*;
use std::io::Write;
use tempfile::NamedTempFile;
fn write_json(content: &str) -> NamedTempFile {
let mut f = NamedTempFile::new().unwrap();
f.write_all(content.as_bytes()).unwrap();
f
}
#[test]
fn test_parse_valid_contracts() {
let json = r#"{
"$schema": "https://raw.githubusercontent.com/cobre-rs/cobre/refs/heads/main/schemas/energy_contracts.schema.json",
"contracts": [
{
"id": 0,
"name": "Importação Argentina",
"operational_start_date": "2024-01-01",
"bus_id": 5,
"type": "import",
"price_per_mwh": 200.0,
"limits": { "min_mw": 0.0, "max_mw": 1000.0 }
},
{
"id": 1,
"name": "Exportação Uruguai",
"operational_start_date": "2024-01-01",
"bus_id": 6,
"type": "export",
"entry_stage_id": 1,
"exit_stage_id": 60,
"price_per_mwh": -150.0,
"limits": { "min_mw": 0.0, "max_mw": 500.0 }
}
]
}"#;
let f = write_json(json);
let contracts = parse_energy_contracts(f.path()).unwrap();
assert_eq!(contracts.len(), 2);
assert_eq!(contracts[0].id, EntityId(0));
assert_eq!(contracts[0].name, "Importação Argentina");
assert_eq!(contracts[0].bus_id, EntityId(5));
assert_eq!(contracts[0].contract_type, ContractType::Import);
assert_eq!(contracts[0].entry_stage_id, None);
assert_eq!(contracts[0].exit_stage_id, None);
assert!((contracts[0].price_per_mwh - 200.0).abs() < f64::EPSILON);
assert!((contracts[0].min_mw - 0.0).abs() < f64::EPSILON);
assert!((contracts[0].max_mw - 1000.0).abs() < f64::EPSILON);
assert_eq!(contracts[1].id, EntityId(1));
assert_eq!(contracts[1].name, "Exportação Uruguai");
assert_eq!(contracts[1].contract_type, ContractType::Export);
assert_eq!(contracts[1].entry_stage_id, Some(1));
assert_eq!(contracts[1].exit_stage_id, Some(60));
assert!((contracts[1].price_per_mwh - (-150.0)).abs() < f64::EPSILON);
assert!((contracts[1].min_mw - 0.0).abs() < f64::EPSILON);
assert!((contracts[1].max_mw - 500.0).abs() < f64::EPSILON);
}
#[test]
fn test_unknown_contract_type() {
let json = r#"{
"contracts": [
{
"id": 0, "name": "Bad", "operational_start_date": "2024-01-01", "bus_id": 0,
"type": "unknown_value",
"price_per_mwh": 100.0,
"limits": { "min_mw": 0.0, "max_mw": 100.0 }
}
]
}"#;
let f = write_json(json);
let err = parse_energy_contracts(f.path()).unwrap_err();
assert!(
matches!(err, LoadError::ParseError { .. }),
"expected ParseError for unknown contract type, got: {err:?}"
);
}
#[test]
fn test_duplicate_contract_id() {
let json = r#"{
"contracts": [
{
"id": 4, "name": "Alpha", "operational_start_date": "2024-01-01", "bus_id": 0,
"type": "import",
"price_per_mwh": 100.0,
"limits": { "min_mw": 0.0, "max_mw": 100.0 }
},
{
"id": 4, "name": "Beta", "operational_start_date": "2024-01-01", "bus_id": 1,
"type": "export",
"price_per_mwh": -50.0,
"limits": { "min_mw": 0.0, "max_mw": 200.0 }
}
]
}"#;
let f = write_json(json);
let err = parse_energy_contracts(f.path()).unwrap_err();
match &err {
LoadError::SchemaError { field, message, .. } => {
assert!(
field.contains("contracts[1].id"),
"field should contain 'contracts[1].id', got: {field}"
);
assert!(
message.contains("duplicate"),
"message should contain 'duplicate', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn test_negative_limits_min_mw() {
let json = r#"{
"contracts": [
{
"id": 0, "name": "Bad", "operational_start_date": "2024-01-01", "bus_id": 0,
"type": "import",
"price_per_mwh": 100.0,
"limits": { "min_mw": -10.0, "max_mw": 100.0 }
}
]
}"#;
let f = write_json(json);
let err = parse_energy_contracts(f.path()).unwrap_err();
match &err {
LoadError::SchemaError { field, message, .. } => {
assert!(
field.contains("limits.min_mw"),
"field should contain 'limits.min_mw', got: {field}"
);
assert!(
message.contains(">= 0.0"),
"message should mention >= 0.0, got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn test_max_mw_less_than_min_mw() {
let json = r#"{
"contracts": [
{
"id": 0, "name": "Bad", "operational_start_date": "2024-01-01", "bus_id": 0,
"type": "import",
"price_per_mwh": 100.0,
"limits": { "min_mw": 500.0, "max_mw": 100.0 }
}
]
}"#;
let f = write_json(json);
let err = parse_energy_contracts(f.path()).unwrap_err();
match &err {
LoadError::SchemaError { field, message, .. } => {
assert!(
field.contains("limits.max_mw"),
"field should contain 'limits.max_mw', got: {field}"
);
assert!(
message.contains("min_mw"),
"message should mention min_mw, got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn test_negative_price_per_mwh_is_valid_for_export() {
let json = r#"{
"contracts": [
{
"id": 0, "name": "Export Revenue", "operational_start_date": "2024-01-01", "bus_id": 0,
"type": "export",
"price_per_mwh": -500.0,
"limits": { "min_mw": 0.0, "max_mw": 200.0 }
}
]
}"#;
let f = write_json(json);
let result = parse_energy_contracts(f.path());
assert!(
result.is_ok(),
"negative price_per_mwh should be valid for export, got: {result:?}"
);
let contracts = result.unwrap();
assert!((contracts[0].price_per_mwh - (-500.0)).abs() < f64::EPSILON);
}
#[test]
fn test_declaration_order_invariance() {
let json_forward = r#"{
"contracts": [
{
"id": 0, "name": "Import A", "operational_start_date": "2024-01-01", "bus_id": 0,
"type": "import",
"price_per_mwh": 100.0,
"limits": { "min_mw": 0.0, "max_mw": 100.0 }
},
{
"id": 1, "name": "Export B", "operational_start_date": "2024-01-01", "bus_id": 1,
"type": "export",
"price_per_mwh": -50.0,
"limits": { "min_mw": 0.0, "max_mw": 200.0 }
}
]
}"#;
let json_reversed = r#"{
"contracts": [
{
"id": 1, "name": "Export B", "operational_start_date": "2024-01-01", "bus_id": 1,
"type": "export",
"price_per_mwh": -50.0,
"limits": { "min_mw": 0.0, "max_mw": 200.0 }
},
{
"id": 0, "name": "Import A", "operational_start_date": "2024-01-01", "bus_id": 0,
"type": "import",
"price_per_mwh": 100.0,
"limits": { "min_mw": 0.0, "max_mw": 100.0 }
}
]
}"#;
let f1 = write_json(json_forward);
let f2 = write_json(json_reversed);
let contracts1 = parse_energy_contracts(f1.path()).unwrap();
let contracts2 = parse_energy_contracts(f2.path()).unwrap();
assert_eq!(
contracts1, contracts2,
"results must be identical regardless of input ordering"
);
assert_eq!(contracts1[0].id, EntityId(0));
assert_eq!(contracts1[1].id, EntityId(1));
}
#[test]
fn test_file_not_found() {
let path = Path::new("/nonexistent/system/energy_contracts.json");
let err = parse_energy_contracts(path).unwrap_err();
match &err {
LoadError::IoError { path: p, .. } => {
assert_eq!(p, path);
}
other => panic!("expected IoError, got: {other:?}"),
}
}
#[test]
fn test_invalid_json() {
let f = write_json(r#"{"contracts": [not valid json}}"#);
let err = parse_energy_contracts(f.path()).unwrap_err();
assert!(
matches!(err, LoadError::ParseError { .. }),
"expected ParseError for invalid JSON, got: {err:?}"
);
}
#[test]
fn test_empty_contracts_array() {
let json = r#"{ "contracts": [] }"#;
let f = write_json(json);
let contracts = parse_energy_contracts(f.path()).unwrap();
assert!(contracts.is_empty());
}
#[test]
fn test_min_equals_max_mw_is_valid() {
let json = r#"{
"contracts": [
{
"id": 0, "name": "Fixed", "operational_start_date": "2024-01-01", "bus_id": 0,
"type": "import",
"price_per_mwh": 100.0,
"limits": { "min_mw": 100.0, "max_mw": 100.0 }
}
]
}"#;
let f = write_json(json);
let result = parse_energy_contracts(f.path());
assert!(
result.is_ok(),
"min_mw == max_mw should be valid, got: {result:?}"
);
}
#[test]
fn test_zero_price_is_valid() {
let json = r#"{
"contracts": [
{
"id": 0, "name": "Free Import", "operational_start_date": "2024-01-01", "bus_id": 0,
"type": "import",
"price_per_mwh": 0.0,
"limits": { "min_mw": 0.0, "max_mw": 100.0 }
}
]
}"#;
let f = write_json(json);
let result = parse_energy_contracts(f.path());
assert!(
result.is_ok(),
"zero price should be valid, got: {result:?}"
);
}
}