use super::order::{Order, OrderKind};
use crate::error::{Error, Result};
use crate::shape::{AsShape, Shape};
use core::hash::{Hash, Hasher};
use core::marker::PhantomData;
use core::panic::{RefUnwindSafe, UnwindSafe};
#[cfg(feature = "serde")]
mod serde;
#[derive(Debug)]
pub(super) struct Layout<T, O>
where
O: Order,
{
major: usize,
minor: usize,
element: PhantomData<T>,
order: PhantomData<O>,
}
impl<T, O> Layout<T, O>
where
O: Order,
{
pub(super) const DEFAULT: Self = Self::new_unchecked(0, 0);
}
impl<T, O> Layout<T, O>
where
O: Order,
{
const fn new(major: usize, minor: usize) -> Result<(Self, usize)> {
let Some(size) = major.checked_mul(minor) else {
return Err(Error::SizeOverflow);
};
match Self::check_size(size) {
Err(error) => Err(error),
Ok(_) => Ok((Self::new_unchecked(major, minor), size)),
}
}
const fn new_unchecked(major: usize, minor: usize) -> Self {
Self {
major,
minor,
element: PhantomData,
order: PhantomData,
}
}
pub(super) fn from_shape<S>(shape: S) -> Result<(Self, usize)>
where
S: AsShape,
{
match O::KIND {
OrderKind::RowMajor => Self::new(shape.nrows(), shape.ncols()),
OrderKind::ColMajor => Self::new(shape.ncols(), shape.nrows()),
}
}
pub(super) fn from_shape_unchecked<S>(shape: S) -> Self
where
S: AsShape,
{
match O::KIND {
OrderKind::RowMajor => Self::new_unchecked(shape.nrows(), shape.ncols()),
OrderKind::ColMajor => Self::new_unchecked(shape.ncols(), shape.nrows()),
}
}
pub(super) const fn to_shape(self) -> Shape {
match O::KIND {
OrderKind::RowMajor => Shape::new(self.major, self.minor),
OrderKind::ColMajor => Shape::new(self.minor, self.major),
}
}
pub(super) const fn with_order<P>(self) -> Layout<T, P>
where
P: Order,
{
Layout::new_unchecked(self.major, self.minor)
}
pub(super) const fn cast<U, P>(self) -> Result<Layout<U, P>>
where
P: Order,
{
match Layout::<U, P>::check_size(self.size()) {
Err(error) => Err(error),
Ok(_) => Ok(Layout::new_unchecked(self.major, self.minor)),
}
}
pub(super) const fn major(&self) -> usize {
self.major
}
pub(super) const fn minor(&self) -> usize {
self.minor
}
pub(super) const fn stride(&self) -> Stride {
Stride(self.minor)
}
pub(super) const fn size(&self) -> usize {
self.major * self.minor
}
pub(super) const fn swap(&mut self) -> &mut Self {
(self.major, self.minor) = (self.minor, self.major);
self
}
}
impl<T, O> Layout<T, O>
where
O: Order,
{
const MAX_SIZE: usize = if size_of::<T>() == 0 {
usize::MAX
} else {
isize::MAX as usize / size_of::<T>()
};
pub(super) const fn check_size(size: usize) -> Result<()> {
if size > Self::MAX_SIZE {
Err(Error::CapacityOverflow)
} else {
Ok(())
}
}
}
unsafe impl<T, O> Send for Layout<T, O> where O: Order {}
unsafe impl<T, O> Sync for Layout<T, O> where O: Order {}
impl<T, O> Unpin for Layout<T, O> where O: Order {}
impl<T, O> UnwindSafe for Layout<T, O> where O: Order {}
impl<T, O> RefUnwindSafe for Layout<T, O> where O: Order {}
impl<T, O> Clone for Layout<T, O>
where
O: Order,
{
fn clone(&self) -> Self {
*self
}
}
impl<T, O> Copy for Layout<T, O> where O: Order {}
impl<T, O> Default for Layout<T, O>
where
O: Order,
{
fn default() -> Self {
Self::new_unchecked(0, 0)
}
}
impl<T, O> Hash for Layout<T, O>
where
O: Order,
{
fn hash<H>(&self, state: &mut H)
where
H: Hasher,
{
self.major.hash(state);
self.minor.hash(state);
}
}
impl<L, R, O> PartialEq<Layout<R, O>> for Layout<L, O>
where
O: Order,
{
fn eq(&self, other: &Layout<R, O>) -> bool {
self.major == other.major && self.minor == other.minor
}
}
impl<T, O> Eq for Layout<T, O> where O: Order {}
#[derive(Debug, Clone, Copy, Default, Hash, PartialEq, Eq, PartialOrd, Ord)]
pub(super) struct Stride(usize);
impl Stride {
pub(super) const fn major(&self) -> usize {
self.0
}
pub(super) const fn minor(&self) -> usize {
1
}
}