use std::collections::HashSet;
use std::path::Path;
use cobre_core::{ComputedParameter, EntityId, ParameterKind, ScalarParameter};
use serde::Deserialize;
use crate::LoadError;
#[derive(Deserialize)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub(crate) struct ScalarParametersFile {
#[serde(rename = "$schema", default)]
_schema: Option<String>,
scalar_parameters: Vec<ScalarParameterJsonEntry>,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub(crate) struct ScalarParameterJsonEntry {
id: i32,
name: String,
kind: String,
value: Option<f64>,
values: Option<Vec<(i32, f64)>>,
computed_spec: Option<ComputedParameter>,
block_values: Option<Vec<(i32, i32, f64)>>,
}
pub fn parse_scalar_parameters_json(path: &Path) -> Result<Vec<ScalarParameter>, LoadError> {
let text = std::fs::read_to_string(path).map_err(|e| LoadError::io(path, e))?;
let file: ScalarParametersFile =
serde_json::from_str(&text).map_err(|e| LoadError::parse(path, e.to_string()))?;
let entries = file.scalar_parameters;
let mut seen_ids: HashSet<i32> = HashSet::with_capacity(entries.len());
let mut seen_names: HashSet<String> = HashSet::with_capacity(entries.len());
let mut result: Vec<ScalarParameter> = Vec::with_capacity(entries.len());
for (i, entry) in entries.into_iter().enumerate() {
if !seen_ids.insert(entry.id) {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].id"),
message: format!("duplicate id {}", entry.id),
});
}
if entry.name.is_empty() {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].name"),
message: "name must not be empty".to_string(),
});
}
if entry.name != entry.name.trim() {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].name"),
message: format!(
"name must not have leading or trailing whitespace, got {:?}",
entry.name
),
});
}
if !seen_names.insert(entry.name.clone()) {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].name"),
message: format!("duplicate name {:?}", entry.name),
});
}
let kind = convert_entry_to_kind(i, &entry, path)?;
result.push(ScalarParameter {
id: EntityId(entry.id),
name: entry.name,
kind,
});
}
result.sort_by_key(|p| p.id.0);
Ok(result)
}
pub fn load_scalar_parameters_json(case_dir: &Path) -> Result<Vec<ScalarParameter>, LoadError> {
parse_scalar_parameters_json(&case_dir.join("constraints/generic_parameters.json"))
}
fn convert_entry_to_kind(
i: usize,
entry: &ScalarParameterJsonEntry,
path: &Path,
) -> Result<ParameterKind, LoadError> {
if entry.kind != "per_stage_block" && entry.block_values.is_some() {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].block_values"),
message: format!(
"\"block_values\" is only valid for kind \"per_stage_block\", not {:?}",
entry.kind
),
});
}
match entry.kind.as_str() {
"constant" => convert_constant(i, entry, path),
"per_stage" => convert_per_stage(i, entry, path),
"seasonal" => convert_seasonal(i, entry, path),
"computed" => convert_computed(i, entry, path),
"per_stage_block" => convert_per_stage_block(i, entry, path),
other => Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].kind"),
message: format!(
"unknown kind {other:?}; legal values are: constant, per_stage, seasonal, computed, per_stage_block"
),
}),
}
}
fn convert_constant(
i: usize,
entry: &ScalarParameterJsonEntry,
path: &Path,
) -> Result<ParameterKind, LoadError> {
let value = entry.value.ok_or_else(|| LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].value"),
message: "\"constant\" kind requires a \"value\" field".to_string(),
})?;
if !value.is_finite() {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].value"),
message: format!("value must be finite, got {value}"),
});
}
Ok(ParameterKind::Constant { value })
}
fn convert_per_stage(
i: usize,
entry: &ScalarParameterJsonEntry,
path: &Path,
) -> Result<ParameterKind, LoadError> {
let pairs = entry
.values
.as_deref()
.ok_or_else(|| LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].values"),
message: "\"per_stage\" kind requires a \"values\" field".to_string(),
})?;
if pairs.is_empty() {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].values"),
message: "\"per_stage\" kind requires at least one entry".to_string(),
});
}
let mut sorted: Vec<(i32, f64)> = pairs.to_vec();
sorted.sort_by_key(|(k, _)| *k);
for window in sorted.windows(2) {
if window[0].0 == window[1].0 {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].values"),
message: format!("duplicate stage_id {} in per_stage values", window[0].0),
});
}
}
for (expected, &(actual, _)) in sorted.iter().enumerate() {
let expected_i32 = i32::try_from(expected).unwrap_or(i32::MAX);
if actual != expected_i32 {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].values"),
message: format!(
"per_stage values must have contiguous stage_ids starting at 0; \
expected stage_id {expected_i32} but got {actual}"
),
});
}
}
for &(stage_id, v) in &sorted {
if !v.is_finite() {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].values"),
message: format!("value for stage_id {stage_id} must be finite, got {v}"),
});
}
}
let dense: Vec<f64> = sorted.into_iter().map(|(_, v)| v).collect();
Ok(ParameterKind::PerStage { values: dense })
}
fn convert_seasonal(
i: usize,
entry: &ScalarParameterJsonEntry,
path: &Path,
) -> Result<ParameterKind, LoadError> {
let pairs = entry
.values
.as_deref()
.ok_or_else(|| LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].values"),
message: "\"seasonal\" kind requires a \"values\" field".to_string(),
})?;
if pairs.is_empty() {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].values"),
message: "\"seasonal\" kind requires at least one entry".to_string(),
});
}
let mut seen_seasons: HashSet<i32> = HashSet::new();
for &(season_id, v) in pairs {
if !seen_seasons.insert(season_id) {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].values"),
message: format!("duplicate season_id {season_id} in seasonal values"),
});
}
if !v.is_finite() {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].values"),
message: format!("value for season_id {season_id} must be finite, got {v}"),
});
}
}
Ok(ParameterKind::new_seasonal(pairs.to_vec()))
}
fn convert_computed(
i: usize,
entry: &ScalarParameterJsonEntry,
path: &Path,
) -> Result<ParameterKind, LoadError> {
let computed_spec = entry.computed_spec.ok_or_else(|| LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].computed_spec"),
message: "\"computed\" kind requires a \"computed_spec\" field".to_string(),
})?;
Ok(ParameterKind::Computed { computed_spec })
}
fn convert_per_stage_block(
i: usize,
entry: &ScalarParameterJsonEntry,
path: &Path,
) -> Result<ParameterKind, LoadError> {
let triples = entry
.block_values
.as_deref()
.ok_or_else(|| LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].block_values"),
message: "\"per_stage_block\" kind requires a \"block_values\" field".to_string(),
})?;
if triples.is_empty() {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].block_values"),
message: "\"per_stage_block\" kind requires at least one entry".to_string(),
});
}
let mut sorted: Vec<(i32, i32, f64)> = triples.to_vec();
sorted.sort_by_key(|&(stage_id, block_id, _)| (stage_id, block_id));
for window in sorted.windows(2) {
if window[0].0 == window[1].0 && window[0].1 == window[1].1 {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].block_values"),
message: format!(
"duplicate (stage_id, block_id) pair ({}, {}) in per_stage_block values",
window[0].0, window[0].1
),
});
}
}
for &(stage_id, block_id, v) in &sorted {
if !v.is_finite() {
return Err(LoadError::SchemaError {
path: path.to_path_buf(),
field: format!("scalar_parameters[{i}].block_values"),
message: format!(
"value for (stage_id {stage_id}, block_id {block_id}) must be finite, got {v}"
),
});
}
}
Ok(ParameterKind::PerStageBlock { values: sorted })
}
#[cfg(test)]
#[allow(
clippy::doc_markdown,
clippy::expect_used,
clippy::panic,
clippy::too_many_lines,
clippy::unwrap_used
)]
mod tests {
use std::io::Write;
use super::*;
use cobre_core::EntityId;
use tempfile::NamedTempFile;
fn write_json(content: &str) -> NamedTempFile {
let mut tmp = NamedTempFile::new().expect("tempfile");
tmp.write_all(content.as_bytes()).expect("write JSON");
tmp
}
#[test]
fn scalar_parameters_json_happy_path_all_four_variants() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": "const_p", "kind": "constant", "value": 3.6 },
{ "id": 2, "name": "stage_p", "kind": "per_stage",
"values": [[0, 1.0], [1, 2.0], [2, 3.0]] },
{ "id": 3, "name": "season_p", "kind": "seasonal",
"values": [[2, 0.8], [1, 1.2]] },
{ "id": 4, "name": "computed_p", "kind": "computed",
"computed_spec": { "tag": "equivalent_productivity", "hydro_id": 7 } }
]
}"#;
let tmp = write_json(json);
let params = parse_scalar_parameters_json(tmp.path()).unwrap();
assert_eq!(params.len(), 4);
assert_eq!(params[0].id, EntityId(1));
assert_eq!(params[0].kind, ParameterKind::Constant { value: 3.6 });
assert_eq!(params[1].id, EntityId(2));
assert_eq!(
params[1].kind,
ParameterKind::PerStage {
values: vec![1.0, 2.0, 3.0]
}
);
assert_eq!(params[2].id, EntityId(3));
assert_eq!(
params[2].kind,
ParameterKind::Seasonal {
values: vec![(1, 1.2), (2, 0.8)]
}
);
assert_eq!(params[3].id, EntityId(4));
assert_eq!(
params[3].kind,
ParameterKind::Computed {
computed_spec: ComputedParameter::EquivalentProductivity {
hydro_id: EntityId(7)
}
}
);
}
#[test]
fn scalar_parameters_json_rejects_duplicate_id() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": "a", "kind": "constant", "value": 3.6 },
{ "id": 1, "name": "b", "kind": "constant", "value": 4.0 }
]
}"#;
let tmp = write_json(json);
let err = parse_scalar_parameters_json(tmp.path()).unwrap_err();
match err {
LoadError::SchemaError { message, .. } => {
assert!(
message.contains("duplicate id"),
"message should contain 'duplicate id', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn scalar_parameters_json_rejects_duplicate_name() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": "rho_eq_h1", "kind": "constant", "value": 3.6 },
{ "id": 2, "name": "rho_eq_h1", "kind": "constant", "value": 4.0 }
]
}"#;
let tmp = write_json(json);
let err = parse_scalar_parameters_json(tmp.path()).unwrap_err();
match err {
LoadError::SchemaError { message, .. } => {
assert!(
message.contains("duplicate name"),
"message should contain 'duplicate name', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn scalar_parameters_json_rejects_empty_name() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": "", "kind": "constant", "value": 3.6 }
]
}"#;
let tmp = write_json(json);
let err = parse_scalar_parameters_json(tmp.path()).unwrap_err();
match err {
LoadError::SchemaError { field, message, .. } => {
assert!(
field.contains("name"),
"field should contain 'name', got: {field}"
);
assert!(
message.contains("empty"),
"message should mention 'empty', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn scalar_parameters_json_rejects_whitespace_name() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": " rho ", "kind": "constant", "value": 3.6 }
]
}"#;
let tmp = write_json(json);
let err = parse_scalar_parameters_json(tmp.path()).unwrap_err();
match err {
LoadError::SchemaError { field, .. } => {
assert!(
field.contains("name"),
"field should contain 'name', got: {field}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn scalar_parameters_json_rejects_seasonal_duplicate_keys() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": "s", "kind": "seasonal",
"values": [[1, 0.5], [1, 0.6]] }
]
}"#;
let tmp = write_json(json);
let err = parse_scalar_parameters_json(tmp.path()).unwrap_err();
match err {
LoadError::SchemaError { message, .. } => {
assert!(
message.contains("duplicate season_id"),
"message should contain 'duplicate season_id', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn scalar_parameters_json_rejects_per_stage_non_contiguous_keys() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": "p", "kind": "per_stage",
"values": [[0, 1.0], [2, 2.0]] }
]
}"#;
let tmp = write_json(json);
let err = parse_scalar_parameters_json(tmp.path()).unwrap_err();
match err {
LoadError::SchemaError { message, .. } => {
assert!(
message.contains("contiguous") || message.contains("non-contiguous"),
"message should mention contiguity, got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn scalar_parameters_json_rejects_per_stage_duplicate_keys() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": "p", "kind": "per_stage",
"values": [[0, 1.0], [0, 2.0], [1, 3.0]] }
]
}"#;
let tmp = write_json(json);
let err = parse_scalar_parameters_json(tmp.path()).unwrap_err();
match err {
LoadError::SchemaError { message, .. } => {
assert!(
message.contains("duplicate stage_id"),
"message should contain 'duplicate stage_id', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn scalar_parameters_json_rejects_non_finite_value() {
let result = convert_constant(
0,
&ScalarParameterJsonEntry {
id: 1,
name: "p".to_string(),
kind: "constant".to_string(),
value: Some(f64::NAN),
values: None,
computed_spec: None,
block_values: None,
},
std::path::Path::new("/test.json"),
);
assert!(result.is_err(), "NaN constant value should be rejected");
let err = result.unwrap_err();
match err {
LoadError::SchemaError { message, .. } => {
assert!(
message.contains("finite"),
"message should mention 'finite', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn scalar_parameters_json_rejects_unknown_field() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": "a", "kind": "constant", "value": 3.6,
"values_source": "sidecar" }
]
}"#;
let tmp = write_json(json);
let err = parse_scalar_parameters_json(tmp.path()).unwrap_err();
match err {
LoadError::ParseError { message, .. } => {
assert!(
message.contains("values_source") || message.contains("unknown field"),
"message should mention 'values_source' or 'unknown field', got: {message}"
);
}
other => panic!("expected ParseError for unknown field, got: {other:?}"),
}
}
#[test]
fn scalar_parameters_json_normalizes_declaration_order() {
let json = r#"{
"scalar_parameters": [
{ "id": 3, "name": "third", "kind": "constant", "value": 3.0 },
{ "id": 1, "name": "first", "kind": "constant", "value": 1.0 },
{ "id": 2, "name": "second", "kind": "constant", "value": 2.0 }
]
}"#;
let tmp = write_json(json);
let params = parse_scalar_parameters_json(tmp.path()).unwrap();
assert_eq!(params.len(), 3);
assert_eq!(params[0].id, EntityId(1));
assert_eq!(params[1].id, EntityId(2));
assert_eq!(params[2].id, EntityId(3));
}
#[test]
fn scalar_parameters_json_rejects_per_stage_empty_values() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": "p", "kind": "per_stage", "values": [] }
]
}"#;
let tmp = write_json(json);
let err = parse_scalar_parameters_json(tmp.path()).unwrap_err();
match err {
LoadError::SchemaError { message, .. } => {
assert!(
message.contains("at least one entry"),
"message should mention 'at least one entry', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn scalar_parameters_json_rejects_seasonal_empty_values() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": "s", "kind": "seasonal", "values": [] }
]
}"#;
let tmp = write_json(json);
let err = parse_scalar_parameters_json(tmp.path()).unwrap_err();
match err {
LoadError::SchemaError { message, .. } => {
assert!(
message.contains("at least one entry"),
"message should mention 'at least one entry', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn scalar_parameters_json_parses_per_stage_block() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": "block_p", "kind": "per_stage_block",
"block_values": [[0, 0, 1.0], [0, 1, 2.0]] }
]
}"#;
let tmp = write_json(json);
let params = parse_scalar_parameters_json(tmp.path()).unwrap();
assert_eq!(params.len(), 1);
assert_eq!(params[0].id, EntityId(1));
assert_eq!(
params[0].kind,
ParameterKind::PerStageBlock {
values: vec![(0, 0, 1.0), (0, 1, 2.0)]
}
);
}
#[test]
fn scalar_parameters_json_per_stage_block_sorts_triples() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": "block_p", "kind": "per_stage_block",
"block_values": [[1, 0, 3.0], [0, 1, 2.0], [0, 0, 1.0]] }
]
}"#;
let tmp = write_json(json);
let params = parse_scalar_parameters_json(tmp.path()).unwrap();
assert_eq!(
params[0].kind,
ParameterKind::PerStageBlock {
values: vec![(0, 0, 1.0), (0, 1, 2.0), (1, 0, 3.0)]
}
);
}
#[test]
fn scalar_parameters_json_rejects_per_stage_block_duplicate_pair() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": "b", "kind": "per_stage_block",
"block_values": [[0, 0, 1.0], [0, 0, 2.0]] }
]
}"#;
let tmp = write_json(json);
let err = parse_scalar_parameters_json(tmp.path()).unwrap_err();
match err {
LoadError::SchemaError { message, .. } => {
assert!(
message.contains("duplicate"),
"message should contain 'duplicate', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn scalar_parameters_json_rejects_block_values_on_wrong_kind() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": "c", "kind": "constant", "value": 3.6,
"block_values": [[0, 0, 1.0]] }
]
}"#;
let tmp = write_json(json);
let err = parse_scalar_parameters_json(tmp.path()).unwrap_err();
match err {
LoadError::SchemaError { field, message, .. } => {
assert!(
field.contains("block_values"),
"field should contain 'block_values', got: {field}"
);
assert!(
message.contains("per_stage_block"),
"message should mention 'per_stage_block', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
#[test]
fn scalar_parameters_json_rejects_per_stage_block_missing_block_values() {
let json = r#"{
"scalar_parameters": [
{ "id": 1, "name": "b", "kind": "per_stage_block" }
]
}"#;
let tmp = write_json(json);
let err = parse_scalar_parameters_json(tmp.path()).unwrap_err();
match err {
LoadError::SchemaError { field, message, .. } => {
assert!(
field.contains("block_values"),
"field should contain 'block_values', got: {field}"
);
assert!(
message.contains("requires"),
"message should mention 'requires', got: {message}"
);
}
other => panic!("expected SchemaError, got: {other:?}"),
}
}
}