use std::ops::{Index, IndexMut};
use crate::errors::QlResult;
use crate::math::array::Array;
use crate::math::timegrid::TimeGrid;
use crate::methods::montecarlo::Path;
use crate::require;
use crate::types::Size;
#[derive(Clone, Debug, PartialEq, Default)]
pub struct MultiPath {
paths: Vec<Path>,
}
impl MultiPath {
pub fn new(n_asset: Size, time_grid: &TimeGrid) -> QlResult<Self> {
require!(n_asset > 0, "number of asset must be positive");
let paths = (0..n_asset)
.map(|_| Path::new(time_grid.clone(), Array::new()))
.collect::<QlResult<Vec<Path>>>()?;
Ok(MultiPath { paths })
}
pub fn from_paths(paths: Vec<Path>) -> Self {
MultiPath { paths }
}
pub fn asset_number(&self) -> Size {
self.paths.len()
}
pub fn path_size(&self) -> Size {
self.paths[0].length()
}
pub fn at(&self, j: Size) -> Option<&Path> {
self.paths.get(j)
}
pub fn at_mut(&mut self, j: Size) -> Option<&mut Path> {
self.paths.get_mut(j)
}
}
impl Index<Size> for MultiPath {
type Output = Path;
fn index(&self, j: Size) -> &Path {
&self.paths[j]
}
}
impl IndexMut<Size> for MultiPath {
fn index_mut(&mut self, j: Size) -> &mut Path {
&mut self.paths[j]
}
}
#[cfg(test)]
mod tests {
use super::*;
fn grid() -> TimeGrid {
TimeGrid::new(1.0, 4).unwrap()
}
#[test]
fn new_builds_n_asset_grid_sized_paths() {
let mp = MultiPath::new(3, &grid()).unwrap();
assert_eq!(mp.asset_number(), 3);
assert_eq!(mp.path_size(), 5);
for j in 0..mp.asset_number() {
assert_eq!(mp[j].values(), &[0.0; 5]);
}
}
#[test]
fn zero_assets_is_rejected() {
let err = MultiPath::new(0, &grid()).unwrap_err();
assert_eq!(err.message(), "number of asset must be positive");
}
#[test]
fn from_paths_round_trips() {
let paths = vec![
Path::new(grid(), Array::from([1.0, 2.0, 3.0, 4.0, 5.0])).unwrap(),
Path::new(grid(), Array::from([6.0, 7.0, 8.0, 9.0, 10.0])).unwrap(),
];
let mp = MultiPath::from_paths(paths);
assert_eq!(mp.asset_number(), 2);
assert_eq!(mp[0].front(), 1.0);
assert_eq!(mp[1].back(), 10.0);
}
#[test]
fn index_and_at_accessors() {
let mut mp = MultiPath::new(2, &grid()).unwrap();
*mp[0].front_mut() = 100.0;
assert_eq!(mp[0].front(), 100.0);
assert!(mp.at(1).is_some());
assert!(mp.at(2).is_none());
*mp.at_mut(1).unwrap().back_mut() = 200.0;
assert_eq!(mp[1].back(), 200.0);
assert!(mp.at_mut(2).is_none());
}
}