1use crate::element_op::Identity;
9use crate::view::validate_bounds;
10use core::marker::PhantomData;
11use core::mem::MaybeUninit;
12use core::ptr::NonNull;
13use num_complex::{Complex32, Complex64};
14
15mod private {
16 pub trait Sealed {}
17}
18
19pub trait KernelStorageElement: private::Sealed + Copy + 'static {
21 const DTYPE: KernelDType;
22}
23
24macro_rules! kernel_storage_element {
25 ($($ty:ty => $dtype:ident),* $(,)?) => {$ (
26 impl private::Sealed for $ty {}
27 impl KernelStorageElement for $ty {
28 const DTYPE: KernelDType = KernelDType::$dtype;
29 }
30 )* };
31}
32
33kernel_storage_element! {
34 f32 => F32,
35 f64 => F64,
36 i32 => I32,
37 i64 => I64,
38 bool => Bool,
39 Complex32 => C32,
40 Complex64 => C64,
41}
42
43use crate::{Result, StridedError, StridedView, StridedViewMut};
44
45#[non_exhaustive]
51#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
52#[repr(u8)]
53pub enum KernelDType {
54 F32 = 1,
55 F64 = 2,
56 I32 = 3,
57 I64 = 4,
58 Bool = 5,
59 C32 = 6,
60 C64 = 7,
61}
62
63impl KernelDType {
64 #[inline]
65 pub const fn label(self) -> &'static str {
66 match self {
67 Self::F32 => "f32",
68 Self::F64 => "f64",
69 Self::I32 => "i32",
70 Self::I64 => "i64",
71 Self::Bool => "bool",
72 Self::C32 => "c32",
73 Self::C64 => "c64",
74 }
75 }
76
77 #[inline]
78 pub const fn size_of(self) -> usize {
79 match self {
80 Self::F32 => core::mem::size_of::<f32>(),
81 Self::F64 => core::mem::size_of::<f64>(),
82 Self::I32 => core::mem::size_of::<i32>(),
83 Self::I64 => core::mem::size_of::<i64>(),
84 Self::Bool => core::mem::size_of::<bool>(),
85 Self::C32 => core::mem::size_of::<Complex32>(),
86 Self::C64 => core::mem::size_of::<Complex64>(),
87 }
88 }
89
90 #[inline]
91 pub const fn alignment(self) -> usize {
92 match self {
93 Self::F32 => core::mem::align_of::<f32>(),
94 Self::F64 => core::mem::align_of::<f64>(),
95 Self::I32 => core::mem::align_of::<i32>(),
96 Self::I64 => core::mem::align_of::<i64>(),
97 Self::Bool => core::mem::align_of::<bool>(),
98 Self::C32 => core::mem::align_of::<Complex32>(),
99 Self::C64 => core::mem::align_of::<Complex64>(),
100 }
101 }
102
103 #[inline]
104 pub const fn requires_valid_byte_values(self) -> bool {
105 matches!(self, Self::Bool)
106 }
107}
108
109fn validate_erased_buffer(dtype: KernelDType, data: &[u8]) -> Result<usize> {
110 let element_count = validate_erased_buffer_extent(dtype, data)?;
111 if element_count != 0 && dtype.requires_valid_byte_values() {
112 validate_bool_bytes(data)?;
113 }
114 Ok(element_count)
115}
116
117fn validate_erased_buffer_extent(dtype: KernelDType, data: &[u8]) -> Result<usize> {
122 let element_size = dtype.size_of();
123 if data.len() % element_size != 0 {
124 return Err(StridedError::ByteLengthMismatch {
125 dtype: dtype.label(),
126 byte_len: data.len(),
127 element_size,
128 });
129 }
130
131 let element_count = data.len() / element_size;
132 if element_count == 0 {
133 return Ok(0);
134 }
135
136 let alignment = dtype.alignment();
137 if data.as_ptr() as usize % alignment != 0 {
138 return Err(StridedError::DataAlignmentMismatch {
139 dtype: dtype.label(),
140 alignment,
141 });
142 }
143 Ok(element_count)
144}
145
146fn validate_bool_bytes(data: &[u8]) -> Result<()> {
152 if data.iter().fold(0u8, |acc, &value| acc | value) <= 1 {
153 return Ok(());
154 }
155 match data.iter().find(|&&value| value > 1) {
156 Some(&value) => Err(StridedError::InvalidBoolByte { value }),
157 None => Ok(()),
158 }
159}
160
161fn validate_erased_buffer_layout(
162 dtype: KernelDType,
163 data: NonNull<u8>,
164 byte_len: usize,
165) -> Result<usize> {
166 let element_size = dtype.size_of();
167 if byte_len % element_size != 0 {
168 return Err(StridedError::ByteLengthMismatch {
169 dtype: dtype.label(),
170 byte_len,
171 element_size,
172 });
173 }
174 if byte_len != 0 && data.as_ptr() as usize % dtype.alignment() != 0 {
175 return Err(StridedError::DataAlignmentMismatch {
176 dtype: dtype.label(),
177 alignment: dtype.alignment(),
178 });
179 }
180 Ok(byte_len / element_size)
181}
182
183#[derive(Clone, Copy, Debug)]
189pub struct ErasedRawStridedPtr<'a> {
190 dtype: KernelDType,
191 data: NonNull<u8>,
192 byte_len: usize,
193 dims: &'a [usize],
194 strides: &'a [isize],
195 offset: isize,
196 marker: PhantomData<&'a [MaybeUninit<u8>]>,
197}
198
199impl<'a> ErasedRawStridedPtr<'a> {
200 pub unsafe fn from_raw_parts(
217 dtype: KernelDType,
218 data: NonNull<u8>,
219 byte_len: usize,
220 dims: &'a [usize],
221 strides: &'a [isize],
222 offset: isize,
223 ) -> Result<Self> {
224 let element_count = validate_erased_buffer_layout(dtype, data, byte_len)?;
225 validate_bounds(element_count, dims, strides, offset)?;
226 Ok(Self {
227 dtype,
228 data,
229 byte_len,
230 dims,
231 strides,
232 offset,
233 marker: PhantomData,
234 })
235 }
236
237 pub fn from_ref(input: &ErasedRawStridedRef<'a>) -> Self {
239 Self {
240 dtype: input.dtype,
241 data: NonNull::new(input.data.as_ptr().cast_mut()).unwrap_or_else(NonNull::dangling),
242 byte_len: input.data.len(),
243 dims: input.dims,
244 strides: input.strides,
245 offset: input.offset,
246 marker: PhantomData,
247 }
248 }
249
250 pub unsafe fn try_as_ref_after_no_overlap(&self) -> Result<ErasedRawStridedRef<'_>> {
261 let bytes = core::slice::from_raw_parts(self.data.as_ptr(), self.byte_len);
262 validate_erased_buffer(self.dtype, bytes)?;
263 ErasedRawStridedRef::from_raw_parts(
264 self.dtype,
265 self.data,
266 self.byte_len,
267 self.dims,
268 self.strides,
269 self.offset,
270 )
271 }
272
273 #[inline]
274 pub fn dtype(&self) -> KernelDType {
275 self.dtype
276 }
277
278 pub fn overlaps_mut(&self, dest: &ErasedRawStridedMut<'_>) -> Result<bool> {
280 ranges_overlap(
281 self.data.as_ptr() as usize,
282 self.byte_len,
283 dest.data.as_ptr() as usize,
284 dest.data.len(),
285 )
286 }
287
288 pub fn overlaps_uninit_mut(&self, dest: &ErasedRawStridedUninitMut<'_>) -> Result<bool> {
290 ranges_overlap(
291 self.data.as_ptr() as usize,
292 self.byte_len,
293 dest.data.as_ptr() as usize,
294 dest.data.len(),
295 )
296 }
297
298 #[inline]
299 pub fn dims(&self) -> &'a [usize] {
300 self.dims
301 }
302
303 #[inline]
304 pub fn strides(&self) -> &'a [isize] {
305 self.strides
306 }
307
308 #[inline]
309 pub fn offset(&self) -> isize {
310 self.offset
311 }
312}
313
314#[derive(Clone, Copy, Debug)]
318pub struct ErasedRawStridedRef<'a> {
319 dtype: KernelDType,
320 data: &'a [u8],
321 dims: &'a [usize],
322 strides: &'a [isize],
323 offset: isize,
324}
325
326impl<'a> ErasedRawStridedRef<'a> {
327 pub fn from_slice<T: KernelStorageElement>(
329 data: &'a [T],
330 dims: &'a [usize],
331 strides: &'a [isize],
332 offset: isize,
333 ) -> Result<Self> {
334 let dtype = T::DTYPE;
335 let byte_len = data
336 .len()
337 .checked_mul(core::mem::size_of::<T>())
338 .ok_or(StridedError::OffsetOverflow)?;
339 let bytes = unsafe { core::slice::from_raw_parts(data.as_ptr().cast::<u8>(), byte_len) };
340 let element_count = validate_erased_buffer_extent(dtype, bytes)?;
341 validate_bounds(element_count, dims, strides, offset)?;
342 Ok(Self {
343 dtype,
344 data: bytes,
345 dims,
346 strides,
347 offset,
348 })
349 }
350
351 pub unsafe fn from_raw_parts(
361 dtype: KernelDType,
362 data: NonNull<u8>,
363 byte_len: usize,
364 dims: &'a [usize],
365 strides: &'a [isize],
366 offset: isize,
367 ) -> Result<Self> {
368 let count = validate_erased_buffer_layout(dtype, data, byte_len)?;
369 let bytes = core::slice::from_raw_parts(data.as_ptr(), byte_len);
370 validate_bounds(count, dims, strides, offset)?;
371 Ok(Self {
372 dtype,
373 data: bytes,
374 dims,
375 strides,
376 offset,
377 })
378 }
379
380 pub fn data_as<T: KernelStorageElement>(&self) -> Result<&[T]> {
382 if self.dtype != T::DTYPE {
383 return Err(StridedError::DTypeMismatch {
384 expected: T::DTYPE.label(),
385 actual: self.dtype.label(),
386 });
387 }
388 Ok(unsafe {
389 core::slice::from_raw_parts(
390 self.data.as_ptr().cast::<T>(),
391 self.data.len() / core::mem::size_of::<T>(),
392 )
393 })
394 }
395
396 #[inline]
397 pub fn dtype(&self) -> KernelDType {
398 self.dtype
399 }
400
401 #[inline]
402 pub fn dims(&self) -> &'a [usize] {
403 self.dims
404 }
405
406 #[inline]
407 pub fn strides(&self) -> &'a [isize] {
408 self.strides
409 }
410
411 #[inline]
412 pub fn offset(&self) -> isize {
413 self.offset
414 }
415}
416
417#[derive(Debug)]
421pub struct ErasedRawStridedMut<'a> {
422 dtype: KernelDType,
423 data: &'a mut [u8],
424 dims: &'a [usize],
425 strides: &'a [isize],
426 offset: isize,
427}
428
429impl<'a> ErasedRawStridedMut<'a> {
430 pub fn from_slice_mut<T: KernelStorageElement>(
432 data: &'a mut [T],
433 dims: &'a [usize],
434 strides: &'a [isize],
435 offset: isize,
436 ) -> Result<Self> {
437 let dtype = T::DTYPE;
438 let byte_len = data
439 .len()
440 .checked_mul(core::mem::size_of::<T>())
441 .ok_or(StridedError::OffsetOverflow)?;
442 let data =
443 unsafe { core::slice::from_raw_parts_mut(data.as_mut_ptr().cast::<u8>(), byte_len) };
444 let element_count = validate_erased_buffer_extent(dtype, data)?;
445 validate_bounds(element_count, dims, strides, offset)?;
446 Ok(Self {
447 dtype,
448 data,
449 dims,
450 strides,
451 offset,
452 })
453 }
454
455 pub unsafe fn from_raw_parts(
466 dtype: KernelDType,
467 data: NonNull<u8>,
468 byte_len: usize,
469 dims: &'a [usize],
470 strides: &'a [isize],
471 offset: isize,
472 ) -> Result<Self> {
473 let count = validate_erased_buffer_layout(dtype, data, byte_len)?;
474 let data = core::slice::from_raw_parts_mut(data.as_ptr(), byte_len);
475 validate_erased_buffer(dtype, data)?;
476 validate_bounds(count, dims, strides, offset)?;
477 Ok(Self {
478 dtype,
479 data,
480 dims,
481 strides,
482 offset,
483 })
484 }
485
486 pub fn data_as<T: KernelStorageElement>(&self) -> Result<&[T]> {
488 if self.dtype != T::DTYPE {
489 return Err(StridedError::DTypeMismatch {
490 expected: T::DTYPE.label(),
491 actual: self.dtype.label(),
492 });
493 }
494 Ok(unsafe {
495 core::slice::from_raw_parts(
496 self.data.as_ptr().cast::<T>(),
497 self.data.len() / core::mem::size_of::<T>(),
498 )
499 })
500 }
501
502 pub fn data_as_mut<T: KernelStorageElement>(&mut self) -> Result<&mut [T]> {
504 if self.dtype != T::DTYPE {
505 return Err(StridedError::DTypeMismatch {
506 expected: T::DTYPE.label(),
507 actual: self.dtype.label(),
508 });
509 }
510 Ok(unsafe {
511 core::slice::from_raw_parts_mut(
512 self.data.as_mut_ptr().cast::<T>(),
513 self.data.len() / core::mem::size_of::<T>(),
514 )
515 })
516 }
517
518 #[inline]
519 pub fn dtype(&self) -> KernelDType {
520 self.dtype
521 }
522
523 #[inline]
524 pub fn dims(&self) -> &'a [usize] {
525 self.dims
526 }
527
528 #[inline]
529 pub fn strides(&self) -> &'a [isize] {
530 self.strides
531 }
532
533 #[inline]
534 pub fn offset(&self) -> isize {
535 self.offset
536 }
537}
538
539#[derive(Debug)]
546pub struct ErasedRawStridedUninitMut<'a> {
547 dtype: KernelDType,
548 data: &'a mut [MaybeUninit<u8>],
549 dims: &'a [usize],
550 strides: &'a [isize],
551 offset: isize,
552}
553
554impl<'a> ErasedRawStridedUninitMut<'a> {
555 pub fn from_uninit_slice<T: KernelStorageElement>(
557 data: &'a mut [MaybeUninit<T>],
558 dims: &'a [usize],
559 strides: &'a [isize],
560 offset: isize,
561 ) -> Result<Self> {
562 let dtype = T::DTYPE;
563 let byte_len = data
564 .len()
565 .checked_mul(core::mem::size_of::<T>())
566 .ok_or(StridedError::OffsetOverflow)?;
567 let data = unsafe {
568 core::slice::from_raw_parts_mut(data.as_mut_ptr().cast::<MaybeUninit<u8>>(), byte_len)
569 };
570 let data_ptr =
571 NonNull::new(data.as_mut_ptr().cast::<u8>()).unwrap_or_else(NonNull::dangling);
572 let element_count = validate_erased_buffer_layout(dtype, data_ptr, byte_len)?;
573 validate_bounds(element_count, dims, strides, offset)?;
574 Ok(Self {
575 dtype,
576 data,
577 dims,
578 strides,
579 offset,
580 })
581 }
582
583 pub unsafe fn from_raw_parts(
594 dtype: KernelDType,
595 data: NonNull<u8>,
596 byte_len: usize,
597 dims: &'a [usize],
598 strides: &'a [isize],
599 offset: isize,
600 ) -> Result<Self> {
601 let count = validate_erased_buffer_layout(dtype, data, byte_len)?;
602 let data =
603 core::slice::from_raw_parts_mut(data.as_ptr().cast::<MaybeUninit<u8>>(), byte_len);
604 validate_bounds(count, dims, strides, offset)?;
605 Ok(Self {
606 dtype,
607 data,
608 dims,
609 strides,
610 offset,
611 })
612 }
613
614 pub fn data_as_uninit_mut<T: KernelStorageElement>(&mut self) -> Result<&mut [MaybeUninit<T>]> {
616 if self.dtype != T::DTYPE {
617 return Err(StridedError::DTypeMismatch {
618 expected: T::DTYPE.label(),
619 actual: self.dtype.label(),
620 });
621 }
622 Ok(unsafe {
623 core::slice::from_raw_parts_mut(
624 self.data.as_mut_ptr().cast::<MaybeUninit<T>>(),
625 self.data.len() / core::mem::size_of::<T>(),
626 )
627 })
628 }
629
630 #[inline]
631 pub fn dtype(&self) -> KernelDType {
632 self.dtype
633 }
634
635 #[inline]
636 pub fn dims(&self) -> &'a [usize] {
637 self.dims
638 }
639
640 #[inline]
641 pub fn strides(&self) -> &'a [isize] {
642 self.strides
643 }
644
645 #[inline]
646 pub fn offset(&self) -> isize {
647 self.offset
648 }
649}
650
651fn ranges_overlap(a_start: usize, a_len: usize, b_start: usize, b_len: usize) -> Result<bool> {
652 let a_end = a_start
653 .checked_add(a_len)
654 .ok_or(StridedError::OffsetOverflow)?;
655 let b_end = b_start
656 .checked_add(b_len)
657 .ok_or(StridedError::OffsetOverflow)?;
658 Ok(a_start < b_end && b_start < a_end)
659}
660
661#[derive(Clone, Copy, Debug)]
667pub struct RawStridedRef<'a, T> {
668 data: &'a [T],
669 dims: &'a [usize],
670 strides: &'a [isize],
671 offset: isize,
672}
673
674impl<'a, T> RawStridedRef<'a, T> {
675 pub fn new(
677 data: &'a [T],
678 dims: &'a [usize],
679 strides: &'a [isize],
680 offset: isize,
681 ) -> Result<Self> {
682 validate_bounds(data.len(), dims, strides, offset)?;
683 Ok(Self {
684 data,
685 dims,
686 strides,
687 offset,
688 })
689 }
690
691 pub unsafe fn new_unchecked(
697 data: &'a [T],
698 dims: &'a [usize],
699 strides: &'a [isize],
700 offset: isize,
701 ) -> Self {
702 Self {
703 data,
704 dims,
705 strides,
706 offset,
707 }
708 }
709
710 #[inline]
711 pub fn data(&self) -> &'a [T] {
712 self.data
713 }
714
715 #[inline]
716 pub fn dims(&self) -> &'a [usize] {
717 self.dims
718 }
719
720 #[inline]
721 pub fn strides(&self) -> &'a [isize] {
722 self.strides
723 }
724
725 #[inline]
726 pub fn offset(&self) -> isize {
727 self.offset
728 }
729
730 #[inline]
731 pub fn ptr(&self) -> *const T {
732 if self.dims.iter().any(|&dim| dim == 0) {
733 NonNull::<T>::dangling().as_ptr()
734 } else {
735 unsafe { self.data.as_ptr().offset(self.offset) }
736 }
737 }
738
739 #[inline]
744 pub fn as_view(&self) -> StridedView<'a, T, Identity> {
745 unsafe { StridedView::new_unchecked(self.data, self.dims, self.strides, self.offset) }
746 }
747}
748
749#[derive(Debug)]
754pub struct RawStridedMut<'a, T> {
755 data: &'a mut [T],
756 dims: &'a [usize],
757 strides: &'a [isize],
758 offset: isize,
759}
760
761impl<'a, T> RawStridedMut<'a, T> {
762 pub fn new(
764 data: &'a mut [T],
765 dims: &'a [usize],
766 strides: &'a [isize],
767 offset: isize,
768 ) -> Result<Self> {
769 validate_bounds(data.len(), dims, strides, offset)?;
770 Ok(Self {
771 data,
772 dims,
773 strides,
774 offset,
775 })
776 }
777
778 pub unsafe fn new_unchecked(
784 data: &'a mut [T],
785 dims: &'a [usize],
786 strides: &'a [isize],
787 offset: isize,
788 ) -> Self {
789 Self {
790 data,
791 dims,
792 strides,
793 offset,
794 }
795 }
796
797 #[inline]
798 pub fn data(&self) -> &[T] {
799 self.data
800 }
801
802 #[inline]
803 pub fn data_mut(&mut self) -> &mut [T] {
804 self.data
805 }
806
807 #[inline]
808 pub fn dims(&self) -> &'a [usize] {
809 self.dims
810 }
811
812 #[inline]
813 pub fn strides(&self) -> &'a [isize] {
814 self.strides
815 }
816
817 #[inline]
818 pub fn offset(&self) -> isize {
819 self.offset
820 }
821
822 #[inline]
823 pub fn ptr(&self) -> *const T {
824 if self.dims.iter().any(|&dim| dim == 0) {
825 NonNull::<T>::dangling().as_ptr()
826 } else {
827 unsafe { self.data.as_ptr().offset(self.offset) }
828 }
829 }
830
831 #[inline]
832 pub fn as_mut_ptr(&mut self) -> *mut T {
833 if self.dims.iter().any(|&dim| dim == 0) {
834 NonNull::<T>::dangling().as_ptr()
835 } else {
836 unsafe { self.data.as_mut_ptr().offset(self.offset) }
837 }
838 }
839
840 #[inline]
845 pub fn as_view(&self) -> StridedView<'_, T, Identity> {
846 unsafe { StridedView::new_unchecked(self.data, self.dims, self.strides, self.offset) }
847 }
848
849 #[inline]
854 pub fn as_view_mut(&mut self) -> StridedViewMut<'_, T> {
855 unsafe { StridedViewMut::new_unchecked(self.data, self.dims, self.strides, self.offset) }
856 }
857}
858
859#[cfg(test)]
860mod tests {
861 use super::*;
862
863 #[test]
864 fn typed_erased_storage_covers_all_kernel_dtypes() {
865 macro_rules! check {
866 ($ty:ty, $value:expr) => {{
867 let mut initialized = [$value, $value];
868 let dims = [2usize];
869 let strides = [1isize];
870 let reference =
871 ErasedRawStridedRef::from_slice(&initialized, &dims, &strides, 0).unwrap();
872 assert_eq!(reference.data_as::<$ty>().unwrap().len(), 2);
873 let mut mutable =
874 ErasedRawStridedMut::from_slice_mut(&mut initialized, &dims, &strides, 0)
875 .unwrap();
876 assert_eq!(mutable.data_as::<$ty>().unwrap().len(), 2);
877 assert_eq!(mutable.data_as_mut::<$ty>().unwrap().len(), 2);
878 let mut uninitialized = [MaybeUninit::<$ty>::uninit(); 2];
879 let mut destination = ErasedRawStridedUninitMut::from_uninit_slice(
880 &mut uninitialized,
881 &dims,
882 &strides,
883 0,
884 )
885 .unwrap();
886 assert_eq!(destination.data_as_uninit_mut::<$ty>().unwrap().len(), 2);
887 }};
888 }
889
890 check!(f32, 1.0f32);
891 check!(f64, 1.0f64);
892 check!(i32, 1i32);
893 check!(i64, 1i64);
894 check!(bool, true);
895 check!(Complex32, Complex32::new(1.0, 0.0));
896 check!(Complex64, Complex64::new(1.0, 0.0));
897 }
898
899 #[test]
900 fn raw_bool_descriptor_reports_the_first_invalid_byte() {
901 let mut bytes = vec![1u8; 4096];
902 bytes[4000] = 3;
903 bytes[4001] = 9;
904 let dims = [bytes.len()];
905 let strides = [1isize];
906 let ptr = NonNull::new(bytes.as_mut_ptr()).unwrap();
907 let erased = unsafe {
908 ErasedRawStridedPtr::from_raw_parts(
909 KernelDType::Bool,
910 ptr,
911 bytes.len(),
912 &dims,
913 &strides,
914 0,
915 )
916 }
917 .unwrap();
918 assert!(matches!(
919 unsafe { erased.try_as_ref_after_no_overlap() },
920 Err(crate::StridedError::InvalidBoolByte { value: 3 })
921 ));
922
923 unsafe {
926 *ptr.as_ptr().add(4000) = 0;
927 *ptr.as_ptr().add(4001) = 1;
928 }
929 assert!(unsafe { erased.try_as_ref_after_no_overlap() }.is_ok());
930 }
931
932 #[test]
933 fn rejects_usize_max_dimension_before_isize_truncation() {
934 let data = [false; 1];
935 let dims = [usize::MAX];
936 let strides = [-1isize];
937 assert!(matches!(
938 RawStridedRef::new(&data, &dims, &strides, 0),
939 Err(crate::StridedError::OffsetOverflow)
940 ));
941 let mut data = [false; 1];
942 assert!(matches!(
943 RawStridedMut::new(&mut data, &dims, &strides, 0),
944 Err(crate::StridedError::OffsetOverflow)
945 ));
946 assert!(matches!(
947 ErasedRawStridedRef::from_slice(&data, &dims, &strides, 0),
948 Err(crate::StridedError::OffsetOverflow)
949 ));
950 let mut data = [false; 1];
951 assert!(matches!(
952 ErasedRawStridedMut::from_slice_mut(&mut data, &dims, &strides, 0),
953 Err(crate::StridedError::OffsetOverflow)
954 ));
955 let ptr = NonNull::new(data.as_mut_ptr()).unwrap();
956 assert!(matches!(
957 unsafe {
958 ErasedRawStridedPtr::from_raw_parts(
959 KernelDType::Bool,
960 ptr.cast(),
961 data.len(),
962 &dims,
963 &strides,
964 0,
965 )
966 },
967 Err(crate::StridedError::OffsetOverflow)
968 ));
969 let mut data = vec![MaybeUninit::<bool>::uninit(); 1];
970 assert!(matches!(
971 ErasedRawStridedUninitMut::from_uninit_slice(&mut data, &dims, &strides, 0),
972 Err(crate::StridedError::OffsetOverflow)
973 ));
974 }
975
976 #[test]
977 fn empty_raw_views_use_dangling_base_pointers_for_extreme_offsets() {
978 let dims = [0usize];
979 let strides = [1isize];
980 let data: [f64; 0] = [];
981 let raw = RawStridedRef::new(&data, &dims, &strides, isize::MAX).unwrap();
982 assert_eq!(raw.ptr(), NonNull::<f64>::dangling().as_ptr());
983
984 let mut data: [f64; 0] = [];
985 let mut raw = RawStridedMut::new(&mut data, &dims, &strides, isize::MAX).unwrap();
986 assert_eq!(raw.ptr(), NonNull::<f64>::dangling().as_ptr());
987 assert_eq!(raw.as_mut_ptr(), NonNull::<f64>::dangling().as_ptr());
988 }
989
990 #[test]
991 fn raw_ref_rejects_out_of_bounds_layout() {
992 let data = [0.0f64; 4];
993 let err = RawStridedRef::new(&data, &[2, 3], &[3, 1], 0).unwrap_err();
994 assert!(matches!(err, crate::StridedError::OffsetOverflow));
995 }
996
997 #[test]
998 fn raw_mut_can_reborrow_as_view() {
999 let mut data = [1, 2, 3, 4];
1000 let mut raw = RawStridedMut::new(&mut data, &[2, 2], &[2, 1], 0).unwrap();
1001 {
1002 let view = raw.as_view();
1003 assert_eq!(view.dims(), &[2, 2]);
1004 }
1005 let view_mut = raw.as_view_mut();
1006 assert_eq!(view_mut.get(&[1, 1]), 4);
1007 }
1008}