#[cfg(not(feature = "std"))]
use alloc::vec::Vec;
use crate::Asn1Matrix;
#[cfg(feature = "derive")]
use crate::Errorizable;
pub type MatrixResult<T> = core::result::Result<T, MatrixError>;
#[cfg_attr(feature = "derive", derive(Errorizable))]
#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub enum MatrixError {
#[cfg_attr(feature = "derive", error("Asn1Matrix: n MUST be in 1..=255 (got {0})"))]
InvalidN(u8),
#[cfg_attr(
feature = "derive",
error("Asn1Matrix: data length MUST equal n*n (n={n}, len={len})")
)]
LengthMismatch { n: u8, len: usize },
}
crate::impl_error_display!(MatrixError {
InvalidN(n) => "Asn1Matrix: n MUST be in 1..=255 (got {n})",
LengthMismatch { n, len } => "Asn1Matrix: data length MUST equal n*n (n={n}, len={len})",
});
pub trait MatrixLike {
fn n(&self) -> u8;
fn get(&self, r: u8, c: u8) -> u8;
fn set(&mut self, r: u8, c: u8, value: u8);
fn fill(&mut self, value: u8);
fn clear(&mut self) {
self.fill(0);
}
}
pub trait IntoMatrixDyn {
fn into_matrix_dyn(self) -> Result<MatrixDyn, MatrixError>;
}
impl IntoMatrixDyn for MatrixDyn {
fn into_matrix_dyn(self) -> Result<MatrixDyn, MatrixError> {
Ok(self)
}
}
impl<M> IntoMatrixDyn for M
where
M: MatrixLike,
MatrixDyn: TryFrom<M, Error = MatrixError>,
{
fn into_matrix_dyn(self) -> Result<MatrixDyn, MatrixError> {
MatrixDyn::try_from(self)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct MatrixDyn {
n: u8,
data: Vec<u8>,
}
impl Default for MatrixDyn {
fn default() -> Self {
Self { n: 1, data: vec![0u8; 1] }
}
}
impl MatrixDyn {
pub fn from_row_major(n: u8, bytes: Vec<u8>) -> Option<Self> {
if n == 0 {
return None;
}
let n_usize = n as usize;
if bytes.len() == n_usize * n_usize {
Some(Self { n, data: bytes })
} else {
None
}
}
#[inline]
fn idx(&self, r: u8, c: u8) -> Option<usize> {
let n = self.n as usize;
let (ru, cu) = (r as usize, c as usize);
(ru < n && cu < n).then(|| ru * n + cu).filter(|i| *i < self.data.len())
}
pub fn as_bytes(&self) -> &[u8] {
&self.data
}
pub fn as_bytes_mut(&mut self) -> &mut [u8] {
&mut self.data
}
pub fn row(&self, r: u8) -> Option<&[u8]> {
if r >= self.n {
return None;
}
let n = self.n as usize;
let start = r as usize * n;
self.data.get(start..start + n)
}
}
impl MatrixLike for MatrixDyn {
fn n(&self) -> u8 {
self.n
}
fn get(&self, r: u8, c: u8) -> u8 {
self.idx(r, c).map(|i| self.data[i]).unwrap_or(0)
}
fn set(&mut self, r: u8, c: u8, value: u8) {
if let Some(i) = self.idx(r, c) {
self.data[i] = value;
}
}
fn fill(&mut self, value: u8) {
for b in &mut self.data {
*b = value;
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Matrix<const N: usize> {
data: [[u8; N]; N],
}
impl<const N: usize> Default for Matrix<N> {
fn default() -> Self {
const { Self::VALID_N };
Self { data: [[0u8; N]; N] }
}
}
impl<const N: usize> Matrix<N> {
const VALID_N: () = assert!(N >= 1 && N <= 255, "Matrix dimension must be 1..=255");
pub fn new() -> Self {
Self::default()
}
pub fn from_row_major(bytes: &[u8]) -> Self {
let mut m = Self::default();
let mut i = 0usize;
for r in 0..N {
for c in 0..N {
if i < bytes.len() {
m.data[r][c] = bytes[i];
}
i += 1;
}
}
m
}
pub fn row(&self, r: u8) -> Option<&[u8; N]> {
if (r as usize) < N {
Some(&self.data[r as usize])
} else {
None
}
}
}
impl<const N: usize> MatrixLike for Matrix<N> {
fn n(&self) -> u8 {
N as u8
}
fn get(&self, r: u8, c: u8) -> u8 {
if (r as usize) < N && (c as usize) < N {
self.data[r as usize][c as usize]
} else {
0
}
}
fn set(&mut self, r: u8, c: u8, value: u8) {
if (r as usize) < N && (c as usize) < N {
self.data[r as usize][c as usize] = value;
}
}
fn fill(&mut self, value: u8) {
for r in 0..N {
for c in 0..N {
self.data[r][c] = value;
}
}
}
}
macro_rules! validate_n {
($n:expr) => {
if $n == 0 {
return Err(MatrixError::InvalidN($n));
}
};
}
impl TryFrom<u8> for MatrixDyn {
type Error = MatrixError;
fn try_from(n: u8) -> Result<Self, Self::Error> {
validate_n!(n);
let n_usize = n as usize;
let data = vec![0u8; n_usize * n_usize];
Ok(Self { n, data })
}
}
impl TryFrom<MatrixDyn> for Asn1Matrix {
type Error = crate::matrix::MatrixError;
fn try_from(mut matrix: MatrixDyn) -> Result<Self, Self::Error> {
let n = matrix.n();
validate_n!(n);
let expected_len = (n as usize) * (n as usize);
if matrix.data.len() != expected_len {
return Err(crate::matrix::MatrixError::LengthMismatch { n, len: matrix.data.len() });
}
let data = core::mem::take(&mut matrix.data);
Ok(Self { n, data })
}
}
macro_rules! asn1_to_matrix_dyn_impl {
($matrix:expr) => {{
let n = $matrix.n;
validate_n!(n);
let n_u8 = n as u8;
let n2 = n_u8 as usize * n_u8 as usize;
if $matrix.data.len() != n2 {
return Err(MatrixError::LengthMismatch { n, len: $matrix.data.len() });
}
}};
}
impl TryFrom<Asn1Matrix> for MatrixDyn {
type Error = MatrixError;
fn try_from(m: Asn1Matrix) -> Result<Self, Self::Error> {
asn1_to_matrix_dyn_impl!(m);
let mut m = m;
let data = core::mem::take(&mut m.data);
let len = data.len();
MatrixDyn::from_row_major(m.n, data).ok_or(MatrixError::LengthMismatch { n: m.n, len })
}
}
impl TryFrom<&Asn1Matrix> for MatrixDyn {
type Error = MatrixError;
fn try_from(m: &Asn1Matrix) -> Result<Self, Self::Error> {
asn1_to_matrix_dyn_impl!(m);
MatrixDyn::from_row_major(m.n, m.data.clone()).ok_or(MatrixError::LengthMismatch { n: m.n, len: m.data.len() })
}
}
impl TryFrom<Option<Asn1Matrix>> for MatrixDyn {
type Error = MatrixError;
fn try_from(m: Option<Asn1Matrix>) -> Result<Self, Self::Error> {
match m {
Some(m) => Self::try_from(m),
None => Ok(Default::default()),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[cfg(feature = "std")]
fn default_matrix_dyn_upholds_invariant() {
let matrix = MatrixDyn::default();
assert_eq!(matrix.n(), 1);
assert_eq!(matrix.as_bytes(), &[0u8]);
assert_eq!(matrix.get(0, 0), 0);
assert!(matrix.row(0).is_some());
}
#[test]
#[cfg(feature = "std")]
fn absent_asn1_matrix_converts_to_valid_default() -> crate::error::Result<()> {
let matrix = MatrixDyn::try_from(None::<Asn1Matrix>)?;
assert_eq!(matrix.n(), 1);
assert_eq!(matrix.get(0, 0), 0);
assert!(matrix.row(0).is_some());
Ok(())
}
#[test]
#[cfg(feature = "std")]
fn test_matrix_like_reality_end_to_end() -> crate::error::Result<()> {
let bytes: Vec<u8> = (0u8..9u8).collect();
let mut dyn_m = MatrixDyn::from_row_major(3, bytes.clone()).expect("n*n bytes");
let mut stat_m: Matrix<3> = Matrix::<3>::from_row_major(&bytes);
fn paint_diag<M: MatrixLike>(m: &mut M) {
let n = m.n();
m.clear();
for i in 0..n {
m.set(i, i, 1);
}
}
assert_eq!(dyn_m.n(), 3);
assert_eq!(stat_m.n(), 3);
assert_eq!(
dyn_m.row(1).ok_or(crate::testing::error::TestingError::InvariantViolated)?,
&[3, 4, 5]
);
assert_eq!(
stat_m
.row(1)
.ok_or(crate::testing::error::TestingError::InvariantViolated)?
.as_slice(),
&[3, 4, 5]
);
assert_eq!(dyn_m.get(2, 2), 8);
assert_eq!(stat_m.get(0, 2), 2);
paint_diag(&mut dyn_m);
paint_diag(&mut stat_m);
for r in 0..3 {
for c in 0..3 {
let dv = dyn_m.get(r, c);
let sv = stat_m.get(r, c);
if r == c {
assert_eq!(dv, 1);
assert_eq!(sv, 1);
} else {
assert_eq!(dv, 0);
assert_eq!(sv, 0);
}
}
}
dyn_m.fill(7);
for r in 0..3 {
for c in 0..3 {
assert_eq!(dyn_m.get(r, c), 7);
}
}
dyn_m.clear();
for r in 0..3 {
for c in 0..3 {
assert_eq!(dyn_m.get(r, c), 0);
}
}
assert_eq!(dyn_m.as_bytes().len(), 9);
let stat_bytes = Matrix::<3>::from_row_major(&bytes);
assert_eq!(
stat_bytes
.row(0)
.ok_or(crate::testing::error::TestingError::InvariantViolated)?,
&[0, 1, 2]
);
assert!(MatrixDyn::from_row_major(3, vec![0u8; 8]).is_none());
Ok(())
}
#[test]
#[cfg(feature = "std")]
fn test_matrix_specification_compliance() -> crate::error::Result<()> {
let test_matrices = vec![
(1u8, vec![42u8]),
(2u8, vec![1, 2, 3, 4]),
(3u8, (0u8..9u8).collect()),
(255u8, vec![0u8; 255 * 255]),
];
for (n, data) in &test_matrices {
let asn1_matrix = Asn1Matrix { n: *n, data: data.clone() };
assert!(*n >= 1);
let expected_len = (*n as usize) * (*n as usize);
assert_eq!(data.len(), expected_len);
let matrix_dyn = MatrixDyn::try_from(&asn1_matrix)?;
assert_eq!(matrix_dyn.n(), *n);
assert_eq!(matrix_dyn.as_bytes(), data.as_slice());
}
let invalid_n = MatrixDyn::try_from(0u8);
assert!(invalid_n.is_err());
let invalid_asn1 = Asn1Matrix { n: 2, data: vec![1, 2, 3] }; let invalid_conversion = MatrixDyn::try_from(invalid_asn1);
assert!(invalid_conversion.is_err());
let matrix_3x3 = MatrixDyn::try_from(3u8)?;
let mut test_matrix = matrix_3x3;
for r in 0..3 {
for c in 0..3 {
let value = (r + 1) * 10 + c;
test_matrix.set(r, c, value);
}
}
let expected_bytes = vec![10, 11, 12, 20, 21, 22, 30, 31, 32];
assert_eq!(test_matrix.as_bytes(), expected_bytes.as_slice());
let mut matrix = MatrixDyn::try_from(2u8)?;
matrix.set(0, 1, 99);
assert_eq!(matrix.get(0, 1), 99);
assert_eq!(matrix.get(5, 5), 0);
matrix.set(5, 5, 123);
matrix.fill(77);
for r in 0..2 {
for c in 0..2 {
assert_eq!(matrix.get(r, c), 77);
}
}
matrix.clear();
for r in 0..2 {
for c in 0..2 {
assert_eq!(matrix.get(r, c), 0);
}
}
let static_matrix: Matrix<3> = Matrix::from_row_major(&[1, 2, 3, 4, 5, 6, 7, 8, 9]);
let dynamic_matrix = MatrixDyn::from_row_major(3, vec![1, 2, 3, 4, 5, 6, 7, 8, 9]).unwrap();
for r in 0..3 {
for c in 0..3 {
assert_eq!(static_matrix.get(r, c), dynamic_matrix.get(r, c));
}
}
Ok(())
}
}