use std::collections::HashMap;
use crate::{Node, PolicyGraphType, SeasonMap, Transition};
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct HorizonGraph {
pub graph_type: PolicyGraphType,
pub annual_discount_rate: f64,
pub transitions: Vec<Transition>,
pub nodes: Vec<Node>,
pub stage_discount_rate_overrides: HashMap<i32, f64>,
pub season_map: Option<SeasonMap>,
}
impl Default for HorizonGraph {
fn default() -> Self {
Self {
graph_type: PolicyGraphType::FiniteHorizon,
annual_discount_rate: 0.0,
transitions: Vec::new(),
nodes: Vec::new(),
stage_discount_rate_overrides: HashMap::new(),
season_map: None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_horizon_graph_construction() {
let transitions = vec![
Transition {
source_id: 1,
target_id: 2,
probability: 1.0,
annual_discount_rate_override: None,
},
Transition {
source_id: 2,
target_id: 3,
probability: 1.0,
annual_discount_rate_override: Some(0.08),
},
Transition {
source_id: 3,
target_id: 4,
probability: 1.0,
annual_discount_rate_override: None,
},
];
let graph = HorizonGraph {
stage_discount_rate_overrides: std::collections::HashMap::new(),
graph_type: PolicyGraphType::FiniteHorizon,
annual_discount_rate: 0.06,
transitions,
nodes: Vec::new(),
season_map: None,
};
assert_eq!(graph.graph_type, PolicyGraphType::FiniteHorizon);
assert!((graph.annual_discount_rate - 0.06).abs() < f64::EPSILON);
assert_eq!(graph.transitions.len(), 3);
assert_eq!(
graph.transitions[1].annual_discount_rate_override,
Some(0.08)
);
assert!(graph.season_map.is_none());
assert!(graph.nodes.is_empty());
}
#[test]
fn test_horizon_graph_carries_nodes() {
let graph = HorizonGraph {
stage_discount_rate_overrides: std::collections::HashMap::new(),
graph_type: PolicyGraphType::FiniteHorizon,
annual_discount_rate: 0.0,
transitions: vec![Transition {
source_id: 10,
target_id: 11,
probability: 1.0,
annual_discount_rate_override: None,
}],
nodes: vec![
Node {
id: 10,
stage_id: 0,
scenario_id: Some(3),
label: Some("root".to_string()),
},
Node {
id: 11,
stage_id: 1,
scenario_id: None,
label: None,
},
],
season_map: None,
};
assert_eq!(graph.nodes.len(), 2);
assert_eq!(graph.nodes[0].stage_id, 0);
assert_eq!(graph.nodes[0].scenario_id, Some(3));
assert_eq!(graph.nodes[1].scenario_id, None);
}
#[cfg(feature = "serde")]
#[test]
fn test_horizon_graph_serde_roundtrip() {
let graph = HorizonGraph {
stage_discount_rate_overrides: std::collections::HashMap::new(),
graph_type: PolicyGraphType::FiniteHorizon,
annual_discount_rate: 0.06,
transitions: vec![
Transition {
source_id: 1,
target_id: 2,
probability: 1.0,
annual_discount_rate_override: None,
},
Transition {
source_id: 2,
target_id: 3,
probability: 1.0,
annual_discount_rate_override: None,
},
],
nodes: Vec::new(),
season_map: None,
};
let json = serde_json::to_string(&graph).unwrap();
assert!(
json.contains("\"graph_type\":\"FiniteHorizon\""),
"JSON did not contain expected graph_type: {json}"
);
assert!(
json.contains("\"annual_discount_rate\":0.06"),
"JSON did not contain expected annual_discount_rate: {json}"
);
let deserialized: HorizonGraph = serde_json::from_str(&json).unwrap();
assert_eq!(graph, deserialized);
}
}