use ironlab_ir::{DataId, Grid, NdArray, Values};
use crate::matrix::Matrix;
#[derive(Debug, Clone, PartialEq)]
pub enum GridCoords {
Vector(Vec<f64>),
Matrix(Matrix),
}
impl From<Vec<f64>> for GridCoords {
fn from(values: Vec<f64>) -> Self {
GridCoords::Vector(values)
}
}
impl From<&Vec<f64>> for GridCoords {
fn from(values: &Vec<f64>) -> Self {
GridCoords::Vector(values.clone())
}
}
impl From<&[f64]> for GridCoords {
fn from(values: &[f64]) -> Self {
GridCoords::Vector(values.to_vec())
}
}
impl<const N: usize> From<[f64; N]> for GridCoords {
fn from(values: [f64; N]) -> Self {
GridCoords::Vector(values.to_vec())
}
}
impl<const N: usize> From<&[f64; N]> for GridCoords {
fn from(values: &[f64; N]) -> Self {
GridCoords::Vector(values.to_vec())
}
}
impl From<Matrix> for GridCoords {
fn from(matrix: Matrix) -> Self {
GridCoords::Matrix(matrix)
}
}
impl From<&Matrix> for GridCoords {
fn from(matrix: &Matrix) -> Self {
GridCoords::Matrix(matrix.clone())
}
}
#[derive(Clone, Copy)]
enum Along {
Cols,
Rows,
}
pub(crate) fn store_grid(
fig: &mut ironlab_ir::Figure,
x: GridCoords,
y: GridCoords,
field: &Matrix,
) -> Grid {
match (x, y) {
(GridCoords::Vector(x), GridCoords::Vector(y)) => Grid::Rectilinear {
x: fig.add_data(NdArray::vector(x)),
y: fig.add_data(NdArray::vector(y)),
},
(x, y) => Grid::Curvilinear {
x: store_node_coordinates(fig, x, field, Along::Cols),
y: store_node_coordinates(fig, y, field, Along::Rows),
},
}
}
fn store_node_coordinates(
fig: &mut ironlab_ir::Figure,
coords: GridCoords,
field: &Matrix,
along: Along,
) -> DataId {
let (rows, cols) = (field.rows(), field.cols());
let matrix = match coords {
GridCoords::Matrix(matrix) => matrix,
GridCoords::Vector(v) => match along {
Along::Cols if v.len() == cols => Matrix::from_fn(rows, cols, |_, col| v[col]),
Along::Rows if v.len() == rows => Matrix::from_fn(rows, cols, |row, _| v[row]),
_ => return fig.add_data(NdArray::vector(v)),
},
};
fig.add_data(matrix_array(&matrix))
}
pub(crate) fn matrix_array(matrix: &Matrix) -> NdArray {
NdArray {
shape: vec![matrix.rows(), matrix.cols()],
values: Values::F64(matrix.values().to_vec()),
}
}