use crate::Count;
pub mod ids {
pub const BOOL: u32 = 1;
pub const I8: u32 = 2;
pub const U8: u32 = 3;
pub const I16: u32 = 4;
pub const U16: u32 = 5;
pub const I32: u32 = 6;
pub const U32: u32 = 7;
pub const I64: u32 = 8;
pub const U64: u32 = 9;
pub const I128: u32 = 10;
pub const U128: u32 = 11;
pub const ISIZE: u32 = 12;
pub const USIZE: u32 = 13;
pub const F32: u32 = 14;
pub const F64: u32 = 15;
pub const CHAR: u32 = 16;
pub const USER_BASE: u32 = 1024;
}
#[cfg(target_endian = "big")]
pub(crate) fn wire_elem_size(id: u32) -> usize {
match id {
ids::BOOL | ids::I8 | ids::U8 => 1,
ids::I16 | ids::U16 => 2,
ids::I32 | ids::U32 | ids::F32 | ids::CHAR => 4,
ids::I64 | ids::U64 | ids::ISIZE | ids::USIZE | ids::F64 => 8,
ids::I128 | ids::U128 => 16,
_ => 1,
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct DatatypeRef {
pub id: u32,
pub size: usize,
}
pub unsafe trait Equivalence: Copy + 'static {
fn equivalent_datatype() -> DatatypeRef;
}
macro_rules! impl_equivalence {
($($t:ty => $id:expr),+ $(,)?) => {$(
unsafe impl Equivalence for $t {
fn equivalent_datatype() -> DatatypeRef {
DatatypeRef { id: $id, size: ::core::mem::size_of::<$t>() }
}
}
)+};
}
impl_equivalence! {
bool => ids::BOOL,
i8 => ids::I8,
u8 => ids::U8,
i16 => ids::I16,
u16 => ids::U16,
i32 => ids::I32,
u32 => ids::U32,
i64 => ids::I64,
u64 => ids::U64,
i128 => ids::I128,
u128 => ids::U128,
isize => ids::ISIZE,
usize => ids::USIZE,
f32 => ids::F32,
f64 => ids::F64,
char => ids::CHAR,
}
pub trait Collection {
fn count(&self) -> Count;
fn as_datatype(&self) -> DatatypeRef;
}
pub unsafe trait Buffer: Collection {
fn as_bytes(&self) -> &[u8];
}
pub unsafe trait BufferMut: Collection {
fn as_bytes_mut(&mut self) -> &mut [u8];
fn scatter_from(&mut self, bytes: &[u8]) {
let dst = self.as_bytes_mut();
let n = dst.len().min(bytes.len());
dst[..n].copy_from_slice(&bytes[..n]);
}
}
#[inline]
fn as_byte_slice<T: Equivalence>(s: &[T]) -> &[u8] {
unsafe { core::slice::from_raw_parts(s.as_ptr() as *const u8, core::mem::size_of_val(s)) }
}
#[inline]
fn as_byte_slice_mut<T: Equivalence>(s: &mut [T]) -> &mut [u8] {
unsafe { core::slice::from_raw_parts_mut(s.as_mut_ptr() as *mut u8, core::mem::size_of_val(s)) }
}
impl<T: Equivalence> Collection for T {
fn count(&self) -> Count {
1
}
fn as_datatype(&self) -> DatatypeRef {
T::equivalent_datatype()
}
}
unsafe impl<T: Equivalence> Buffer for T {
fn as_bytes(&self) -> &[u8] {
as_byte_slice(core::slice::from_ref(self))
}
}
unsafe impl<T: Equivalence> BufferMut for T {
fn as_bytes_mut(&mut self) -> &mut [u8] {
as_byte_slice_mut(core::slice::from_mut(self))
}
}
impl<T: Equivalence> Collection for [T] {
fn count(&self) -> Count {
self.len() as Count
}
fn as_datatype(&self) -> DatatypeRef {
T::equivalent_datatype()
}
}
unsafe impl<T: Equivalence> Buffer for [T] {
fn as_bytes(&self) -> &[u8] {
as_byte_slice(self)
}
}
unsafe impl<T: Equivalence> BufferMut for [T] {
fn as_bytes_mut(&mut self) -> &mut [u8] {
as_byte_slice_mut(self)
}
}
impl<T: Equivalence, const N: usize> Collection for [T; N] {
fn count(&self) -> Count {
N as Count
}
fn as_datatype(&self) -> DatatypeRef {
T::equivalent_datatype()
}
}
unsafe impl<T: Equivalence, const N: usize> Buffer for [T; N] {
fn as_bytes(&self) -> &[u8] {
as_byte_slice(self.as_slice())
}
}
unsafe impl<T: Equivalence, const N: usize> BufferMut for [T; N] {
fn as_bytes_mut(&mut self) -> &mut [u8] {
as_byte_slice_mut(self.as_mut_slice())
}
}
pub unsafe trait Partitioned {
fn as_datatype(&self) -> DatatypeRef;
fn counts(&self) -> &[Count];
fn displs(&self) -> &[Count];
}
pub unsafe trait PartitionedBuffer: Partitioned {
fn as_bytes(&self) -> &[u8];
}
pub unsafe trait PartitionedBufferMut: Partitioned {
fn as_bytes_mut(&mut self) -> &mut [u8];
}
pub struct Partition<'a, T: Equivalence> {
buf: &'a [T],
counts: Vec<Count>,
displs: Vec<Count>,
}
impl<'a, T: Equivalence> Partition<'a, T> {
pub fn new<C, D>(buf: &'a [T], counts: C, displs: D) -> Partition<'a, T>
where
C: Into<Vec<Count>>,
D: Into<Vec<Count>>,
{
Partition {
buf,
counts: counts.into(),
displs: displs.into(),
}
}
}
unsafe impl<T: Equivalence> Partitioned for Partition<'_, T> {
fn as_datatype(&self) -> DatatypeRef {
T::equivalent_datatype()
}
fn counts(&self) -> &[Count] {
&self.counts
}
fn displs(&self) -> &[Count] {
&self.displs
}
}
unsafe impl<T: Equivalence> PartitionedBuffer for Partition<'_, T> {
fn as_bytes(&self) -> &[u8] {
as_byte_slice(self.buf)
}
}
pub struct PartitionMut<'a, T: Equivalence> {
buf: &'a mut [T],
counts: Vec<Count>,
displs: Vec<Count>,
}
impl<'a, T: Equivalence> PartitionMut<'a, T> {
pub fn new<C, D>(buf: &'a mut [T], counts: C, displs: D) -> PartitionMut<'a, T>
where
C: Into<Vec<Count>>,
D: Into<Vec<Count>>,
{
PartitionMut {
buf,
counts: counts.into(),
displs: displs.into(),
}
}
}
unsafe impl<T: Equivalence> Partitioned for PartitionMut<'_, T> {
fn as_datatype(&self) -> DatatypeRef {
T::equivalent_datatype()
}
fn counts(&self) -> &[Count] {
&self.counts
}
fn displs(&self) -> &[Count] {
&self.displs
}
}
unsafe impl<T: Equivalence> PartitionedBufferMut for PartitionMut<'_, T> {
fn as_bytes_mut(&mut self) -> &mut [u8] {
as_byte_slice_mut(self.buf)
}
}
#[derive(Clone, Debug)]
pub struct UserDatatype {
id: u32,
base: DatatypeRef,
blocks: Vec<(usize, usize)>,
extent: usize,
base_count: Count,
}
impl UserDatatype {
fn from_blocks(base: DatatypeRef, blocks: Vec<(usize, usize)>, extent: usize) -> UserDatatype {
let bytes: usize = blocks.iter().map(|(_, l)| *l).sum();
let base_count = (bytes / base.size.max(1)) as Count;
UserDatatype {
id: ids::USER_BASE + base.id,
base,
blocks,
extent,
base_count,
}
}
pub fn contiguous(count: Count, base: DatatypeRef) -> UserDatatype {
let len = count as usize * base.size;
UserDatatype::from_blocks(base, vec![(0, len)], len)
}
pub fn vector(
count: Count,
blocklength: Count,
stride: Count,
base: DatatypeRef,
) -> UserDatatype {
let bl = blocklength as usize * base.size;
let mut blocks = Vec::with_capacity(count as usize);
for i in 0..count as usize {
blocks.push((i * stride as usize * base.size, bl));
}
let extent = if count == 0 {
0
} else {
((count as usize - 1) * stride as usize + blocklength as usize) * base.size
};
UserDatatype::from_blocks(base, blocks, extent)
}
pub fn indexed(
blocklengths: &[Count],
displacements: &[Count],
base: DatatypeRef,
) -> UserDatatype {
assert_eq!(blocklengths.len(), displacements.len());
let mut blocks = Vec::with_capacity(blocklengths.len());
let mut extent = 0usize;
for (bl, d) in blocklengths.iter().zip(displacements) {
let off = *d as usize * base.size;
let len = *bl as usize * base.size;
blocks.push((off, len));
extent = extent.max(off + len);
}
UserDatatype::from_blocks(base, blocks, extent)
}
pub fn structured(
blocklengths: &[Count],
byte_displacements: &[isize],
base: DatatypeRef,
) -> UserDatatype {
assert_eq!(blocklengths.len(), byte_displacements.len());
let mut blocks = Vec::with_capacity(blocklengths.len());
let mut extent = 0usize;
for (bl, d) in blocklengths.iter().zip(byte_displacements) {
let off = *d as usize;
let len = *bl as usize * base.size;
blocks.push((off, len));
extent = extent.max(off + len);
}
UserDatatype::from_blocks(base, blocks, extent)
}
pub fn id(&self) -> u32 {
self.id
}
pub fn extent(&self) -> usize {
self.extent
}
pub fn base_count(&self) -> Count {
self.base_count
}
pub fn pack(&self, base: &[u8], count: Count) -> Vec<u8> {
let mut out =
Vec::with_capacity(count as usize * self.base_count as usize * self.base.size);
for e in 0..count as usize {
let elem = e * self.extent;
for &(off, len) in &self.blocks {
out.extend_from_slice(&base[elem + off..elem + off + len]);
}
}
out
}
pub fn unpack(&self, packed: &[u8], base: &mut [u8], count: Count) {
let mut pos = 0usize;
for e in 0..count as usize {
let elem = e * self.extent;
for &(off, len) in &self.blocks {
let end = (pos + len).min(packed.len());
let take = end - pos;
base[elem + off..elem + off + take].copy_from_slice(&packed[pos..pos + take]);
pos += take;
}
}
}
}
pub struct View {
packed: Vec<u8>,
base: DatatypeRef,
}
impl View {
pub fn with_count<B: Buffer + ?Sized>(buf: &B, datatype: &UserDatatype, count: Count) -> View {
View {
packed: datatype.pack(buf.as_bytes(), count),
base: datatype.base,
}
}
}
impl Collection for View {
fn count(&self) -> Count {
(self.packed.len() / self.base.size.max(1)) as Count
}
fn as_datatype(&self) -> DatatypeRef {
self.base
}
}
unsafe impl Buffer for View {
fn as_bytes(&self) -> &[u8] {
&self.packed
}
}
pub struct MutView<'a, B: BufferMut + ?Sized> {
base: &'a mut B,
datatype: UserDatatype,
count: Count,
}
impl<'a, B: BufferMut + ?Sized> MutView<'a, B> {
pub fn with_count(buf: &'a mut B, datatype: UserDatatype, count: Count) -> MutView<'a, B> {
MutView {
base: buf,
datatype,
count,
}
}
}
impl<B: BufferMut + ?Sized> Collection for MutView<'_, B> {
fn count(&self) -> Count {
self.datatype.base_count * self.count
}
fn as_datatype(&self) -> DatatypeRef {
self.datatype.base
}
}
unsafe impl<B: BufferMut + ?Sized> BufferMut for MutView<'_, B> {
fn as_bytes_mut(&mut self) -> &mut [u8] {
self.base.as_bytes_mut()
}
fn scatter_from(&mut self, bytes: &[u8]) {
self.datatype
.unpack(bytes, self.base.as_bytes_mut(), self.count);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn datatype_ids_and_sizes() {
assert_eq!(i32::equivalent_datatype().id, ids::I32);
assert_eq!(i32::equivalent_datatype().size, 4);
assert_eq!(f64::equivalent_datatype().size, 8);
assert_eq!(u8::equivalent_datatype().size, 1);
}
#[test]
fn slice_byte_view_roundtrips() {
let data: [i32; 3] = [1, 2, 0x0102_0304];
let bytes = data.as_slice().as_bytes();
assert_eq!(bytes.len(), 12);
assert_eq!(<[i32] as Collection>::count(&data[..]), 3);
let mut out = [0i32; 3];
out.as_mut_slice().as_bytes_mut().copy_from_slice(bytes);
assert_eq!(out, data);
}
#[test]
fn scalar_is_a_buffer_of_one() {
let x = 7u64;
assert_eq!(Collection::count(&x), 1);
assert_eq!(x.as_bytes().len(), 8);
}
#[test]
fn vector_datatype_pack_unpack() {
let dt = UserDatatype::vector(3, 1, 2, i32::equivalent_datatype());
let base = [10i32, 11, 20, 21, 30, 31];
let packed = dt.pack(base.as_slice().as_bytes(), 1);
let got: Vec<i32> = crate::point_to_point::vec_from_bytes(&packed);
assert_eq!(got, vec![10, 20, 30]);
let mut out = [0i32; 6];
dt.unpack(&packed, out.as_mut_slice().as_bytes_mut(), 1);
assert_eq!(out, [10, 0, 20, 0, 30, 0]);
}
#[test]
fn contiguous_and_indexed() {
let c = UserDatatype::contiguous(4, f64::equivalent_datatype());
assert_eq!(c.base_count(), 4);
assert_eq!(c.extent(), 32);
let idx = UserDatatype::indexed(&[2, 1], &[0, 4], i32::equivalent_datatype());
let base = [1i32, 2, 3, 4, 5, 6];
let packed = idx.pack(base.as_slice().as_bytes(), 1);
let got: Vec<i32> = crate::point_to_point::vec_from_bytes(&packed);
assert_eq!(got, vec![1, 2, 5]);
}
#[test]
fn vector_pack_unpack_roundtrips() {
let mut state: u64 = 0xdead_beef_0000_0001;
let mut rng = || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
state
};
for _ in 0..300 {
let count = (rng() % 4) as i32 + 1;
let blocklen = (rng() % 3) as i32 + 1;
let stride = blocklen + (rng() % 3) as i32; let dt = UserDatatype::vector(count, blocklen, stride, i32::equivalent_datatype());
let n_elems = dt.extent() / 4;
let base: Vec<i32> = (0..n_elems as i32).collect();
let packed = dt.pack(base.as_slice().as_bytes(), 1);
let got: Vec<i32> = crate::point_to_point::vec_from_bytes(&packed);
let mut expected = Vec::new();
for c in 0..count {
for k in 0..blocklen {
expected.push(c * stride + k);
}
}
assert_eq!(got, expected, "vector pack mismatch");
let mut out = vec![0i32; n_elems];
dt.unpack(&packed, out.as_mut_slice().as_bytes_mut(), 1);
for c in 0..count {
for k in 0..blocklen {
let idx = (c * stride + k) as usize;
assert_eq!(out[idx], idx as i32, "vector unpack mismatch");
}
}
}
}
}