use cobre_core::EntityId;
use serde::Deserialize;
use std::collections::HashSet;
use std::path::Path;
use crate::LoadError;
#[derive(Deserialize)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
#[serde(deny_unknown_fields)]
pub(crate) struct RawLoadFactorsFile {
#[serde(rename = "$schema")]
_schema: Option<String>,
load_factors: Vec<RawLoadFactorEntry>,
}
#[derive(Deserialize)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
#[serde(deny_unknown_fields)]
struct RawLoadFactorEntry {
bus_id: i32,
stage_id: i32,
block_factors: Vec<RawBlockFactor>,
}
#[derive(Deserialize)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
#[serde(deny_unknown_fields)]
struct RawBlockFactor {
block_id: i32,
factor: f64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct BlockFactor {
pub block_id: i32,
pub factor: f64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct LoadFactorEntry {
pub bus_id: EntityId,
pub stage_id: i32,
pub block_factors: Vec<BlockFactor>,
}
pub fn parse_load_factors(path: &Path) -> Result<Vec<LoadFactorEntry>, LoadError> {
let raw_text = std::fs::read_to_string(path).map_err(|e| LoadError::io(path, e))?;
let raw: RawLoadFactorsFile =
serde_json::from_str(&raw_text).map_err(|e| LoadError::parse(path, e.to_string()))?;
validate_raw(&raw, path)?;
Ok(convert(raw))
}
fn validate_raw(raw: &RawLoadFactorsFile, path: &Path) -> Result<(), LoadError> {
validate_no_duplicate_entries(&raw.load_factors, path)?;
for (i, entry) in raw.load_factors.iter().enumerate() {
validate_block_factors(&entry.block_factors, i, path)?;
}
Ok(())
}
fn validate_no_duplicate_entries(
entries: &[RawLoadFactorEntry],
path: &Path,
) -> Result<(), LoadError> {
let mut seen: HashSet<(i32, i32)> = HashSet::new();
for (i, entry) in entries.iter().enumerate() {
let key = (entry.bus_id, entry.stage_id);
if !seen.insert(key) {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("load_factors[{i}]"),
message: format!(
"duplicate (bus_id={}, stage_id={}) in load_factors",
entry.bus_id, entry.stage_id
),
});
}
}
Ok(())
}
fn validate_block_factors(
block_factors: &[RawBlockFactor],
entry_idx: usize,
path: &Path,
) -> Result<(), LoadError> {
for (j, bf) in block_factors.iter().enumerate() {
if !bf.factor.is_finite() || bf.factor <= 0.0 {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("load_factors[{entry_idx}].block_factors[{j}].factor"),
message: format!("factor must be finite and > 0.0, got {}", bf.factor),
});
}
}
Ok(())
}
fn convert(raw: RawLoadFactorsFile) -> Vec<LoadFactorEntry> {
let mut entries: Vec<LoadFactorEntry> = raw
.load_factors
.into_iter()
.map(|e| {
let mut block_factors: Vec<BlockFactor> = e
.block_factors
.into_iter()
.map(|bf| BlockFactor {
block_id: bf.block_id,
factor: bf.factor,
})
.collect();
block_factors.sort_by_key(|bf| bf.block_id);
LoadFactorEntry {
bus_id: EntityId::from(e.bus_id),
stage_id: e.stage_id,
block_factors,
}
})
.collect();
entries.sort_by(|a, b| {
a.bus_id
.0
.cmp(&b.bus_id.0)
.then_with(|| a.stage_id.cmp(&b.stage_id))
});
entries
}
#[cfg(test)]
#[allow(
clippy::doc_markdown,
clippy::panic,
clippy::too_many_lines,
clippy::unwrap_used
)]
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
}
const VALID_JSON: &str = r#"{
"load_factors": [
{
"bus_id": 0,
"stage_id": 0,
"block_factors": [
{ "block_id": 0, "factor": 0.85 },
{ "block_id": 1, "factor": 1.0 },
{ "block_id": 2, "factor": 1.15 }
]
},
{
"bus_id": 1,
"stage_id": 0,
"block_factors": [
{ "block_id": 0, "factor": 0.90 },
{ "block_id": 1, "factor": 1.05 },
{ "block_id": 2, "factor": 1.20 }
]
}
]
}"#;
#[test]
fn test_valid_2_entries_sorted_and_block_factors_correct() {
let tmp = write_json(VALID_JSON);
let entries = parse_load_factors(tmp.path()).unwrap();
assert_eq!(entries.len(), 2);
assert_eq!(entries[0].bus_id, EntityId::from(0));
assert_eq!(entries[0].stage_id, 0);
assert_eq!(entries[0].block_factors.len(), 3);
assert_eq!(entries[0].block_factors[0].block_id, 0);
assert!((entries[0].block_factors[0].factor - 0.85).abs() < 1e-10);
assert_eq!(entries[0].block_factors[1].block_id, 1);
assert!((entries[0].block_factors[1].factor - 1.0).abs() < f64::EPSILON);
assert_eq!(entries[0].block_factors[2].block_id, 2);
assert!((entries[0].block_factors[2].factor - 1.15).abs() < 1e-10);
assert_eq!(entries[1].bus_id, EntityId::from(1));
assert_eq!(entries[1].stage_id, 0);
assert_eq!(entries[1].block_factors.len(), 3);
}
#[test]
fn test_entries_sorted_by_bus_stage() {
let json = r#"{
"load_factors": [
{
"bus_id": 1,
"stage_id": 0,
"block_factors": [{ "block_id": 0, "factor": 1.0 }]
},
{
"bus_id": 0,
"stage_id": 1,
"block_factors": [{ "block_id": 0, "factor": 0.9 }]
},
{
"bus_id": 0,
"stage_id": 0,
"block_factors": [{ "block_id": 0, "factor": 0.8 }]
}
]
}"#;
let tmp = write_json(json);
let entries = parse_load_factors(tmp.path()).unwrap();
assert_eq!(entries.len(), 3);
assert_eq!(entries[0].bus_id, EntityId::from(0));
assert_eq!(entries[0].stage_id, 0);
assert_eq!(entries[1].bus_id, EntityId::from(0));
assert_eq!(entries[1].stage_id, 1);
assert_eq!(entries[2].bus_id, EntityId::from(1));
assert_eq!(entries[2].stage_id, 0);
}
#[test]
fn test_block_factors_sorted_by_block_id() {
let json = r#"{
"load_factors": [
{
"bus_id": 0,
"stage_id": 0,
"block_factors": [
{ "block_id": 2, "factor": 1.15 },
{ "block_id": 0, "factor": 0.85 },
{ "block_id": 1, "factor": 1.0 }
]
}
]
}"#;
let tmp = write_json(json);
let entries = parse_load_factors(tmp.path()).unwrap();
assert_eq!(entries.len(), 1);
let bfs = &entries[0].block_factors;
assert_eq!(bfs[0].block_id, 0);
assert_eq!(bfs[1].block_id, 1);
assert_eq!(bfs[2].block_id, 2);
}
#[test]
fn test_zero_factor_rejected() {
let json = r#"{
"load_factors": [
{
"bus_id": 0,
"stage_id": 0,
"block_factors": [{ "block_id": 0, "factor": 0.0 }]
}
]
}"#;
let tmp = write_json(json);
let err = parse_load_factors(tmp.path()).unwrap_err();
match &err {
LoadError::SchemaError { field, .. } => {
assert!(
field.contains("factor"),
"field should contain 'factor', got: {field}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn test_negative_factor_rejected() {
let json = r#"{
"load_factors": [
{
"bus_id": 0,
"stage_id": 0,
"block_factors": [{ "block_id": 0, "factor": -0.5 }]
}
]
}"#;
let tmp = write_json(json);
let err = parse_load_factors(tmp.path()).unwrap_err();
match &err {
LoadError::SchemaError { field, .. } => {
assert!(
field.contains("factor"),
"field should contain 'factor', got: {field}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn test_duplicate_bus_stage_rejected() {
let json = r#"{
"load_factors": [
{
"bus_id": 0,
"stage_id": 0,
"block_factors": [{ "block_id": 0, "factor": 1.0 }]
},
{
"bus_id": 0,
"stage_id": 0,
"block_factors": [{ "block_id": 0, "factor": 1.1 }]
}
]
}"#;
let tmp = write_json(json);
let err = parse_load_factors(tmp.path()).unwrap_err();
match &err {
LoadError::SchemaError { field, message, .. } => {
assert!(
field.contains("load_factors[1]"),
"field should contain 'load_factors[1]', got: {field}"
);
assert!(
message.contains("duplicate"),
"message should contain 'duplicate', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn test_empty_load_factors_returns_empty_vec() {
let json = r#"{ "load_factors": [] }"#;
let tmp = write_json(json);
let entries = parse_load_factors(tmp.path()).unwrap();
assert!(entries.is_empty());
}
#[test]
fn test_missing_block_factors_field_is_parse_error() {
let json = r#"{
"load_factors": [
{ "bus_id": 0, "stage_id": 0 }
]
}"#;
let tmp = write_json(json);
let err = parse_load_factors(tmp.path()).unwrap_err();
match &err {
LoadError::ParseError { .. } => {}
other => panic!("expected ParseError, got: {other:?}"),
}
}
#[test]
fn test_very_small_positive_factor_accepted() {
let json = r#"{
"load_factors": [
{
"bus_id": 0,
"stage_id": 0,
"block_factors": [{ "block_id": 0, "factor": 1e-300 }]
}
]
}"#;
let tmp = write_json(json);
let entries = parse_load_factors(tmp.path()).unwrap();
assert_eq!(entries.len(), 1);
assert!(entries[0].block_factors[0].factor > 0.0);
}
#[test]
fn test_declaration_order_invariance() {
let json_fwd = r#"{
"load_factors": [
{
"bus_id": 0, "stage_id": 0,
"block_factors": [{ "block_id": 0, "factor": 1.0 }]
},
{
"bus_id": 1, "stage_id": 0,
"block_factors": [{ "block_id": 0, "factor": 1.1 }]
}
]
}"#;
let json_rev = r#"{
"load_factors": [
{
"bus_id": 1, "stage_id": 0,
"block_factors": [{ "block_id": 0, "factor": 1.1 }]
},
{
"bus_id": 0, "stage_id": 0,
"block_factors": [{ "block_id": 0, "factor": 1.0 }]
}
]
}"#;
let tmp_fwd = write_json(json_fwd);
let tmp_rev = write_json(json_rev);
let entries_fwd = parse_load_factors(tmp_fwd.path()).unwrap();
let entries_rev = parse_load_factors(tmp_rev.path()).unwrap();
let keys_fwd: Vec<(i32, i32)> = entries_fwd
.iter()
.map(|e| (e.bus_id.0, e.stage_id))
.collect();
let keys_rev: Vec<(i32, i32)> = entries_rev
.iter()
.map(|e| (e.bus_id.0, e.stage_id))
.collect();
assert_eq!(keys_fwd, keys_rev);
}
}