Skip to main content

strided_view/
raw.rs

1//! Borrowed raw strided layout types.
2//!
3//! These are the prepared-replay counterparts to [`crate::StridedView`] and
4//! [`crate::StridedViewMut`]. They borrow shape/stride metadata instead of
5//! owning it, so compiled kernels can reuse already-validated layout
6//! descriptors without rebuilding dynamic-rank view wrappers.
7
8use 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
19/// A scalar type that can safely witness storage for an erased kernel view.
20pub 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/// Dtypes supported by dtype-erased kernel entry points.
46///
47/// The enum is intentionally limited to the scalar set currently used by the
48/// tensor runtime callers. Later FFI layers should map their ABI dtype tags to
49/// this enum before dispatching into prepared kernels.
50#[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
117/// Validate the length and alignment of erased storage, without reading it.
118///
119/// Typed constructors call this alone: `KernelStorageElement` is sealed, so a
120/// `&[T]` already holds valid values of `T::DTYPE`.
121fn 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
146/// Reject any byte other than 0 or 1.
147///
148/// The OR fold has no early exit, so it vectorizes and runs at memory
149/// bandwidth; a byte by byte `find` does not, and cost more than the select
150/// kernel it guarded. The slow scan runs only to report the offending byte.
151fn 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/// Pointer-backed dtype-erased input used by one-shot write entry points.
184///
185/// Unlike [`ErasedRawStridedRef`], this descriptor does not create a shared
186/// Rust reference at construction time. That lets the entry point reject an
187/// input/output overlap before forming references to either side.
188#[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    /// Create a pointer-backed descriptor from erased storage.
201    ///
202    /// # Safety
203    ///
204    /// `data` must point into an allocation whose alignment is suitable for
205    /// `dtype`, independently of the observed address. The allocation must
206    /// provide `byte_len` initialized, readable bytes with valid provenance
207    /// for `'a`, and remain alive for `'a`. For `bool`, the bytes may contain
208    /// invalid values temporarily. They must not be read or used to form a
209    /// typed reference until a caller has rejected overlap and
210    /// [`Self::try_as_ref_after_no_overlap`] has validated the complete extent.
211    /// The allocation may overlap a destination.
212    /// The original owner may perform synchronized sequential mutation through
213    /// the same-provenance raw pointer before conversion. No concurrent
214    /// mutation or conflicting access is permitted during overlap checking,
215    /// conversion, or consumer access.
216    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    /// Borrow a safe erased input as a pointer descriptor.
238    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    /// Convert this pointer to an initialized descriptor after overlap checks.
251    ///
252    /// # Safety
253    ///
254    /// The caller must have proved that no mutable allocation overlaps this
255    /// pointer's complete byte extent.
256    /// The allocation contract from [`Self::from_raw_parts`] must still hold.
257    /// No mutation or other conflicting access may occur during this conversion
258    /// or for the duration of the returned descriptor's consumer access.
259    /// Bool bytes are validated here, immediately before typed access.
260    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    /// Check overlap with initialized mutable storage without reading it.
279    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    /// Check overlap with uninitialized mutable storage without reading it.
289    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/// Borrowed dtype-erased raw strided input layout.
315///
316/// `dims`, `strides`, and `offset` are expressed in dtype elements, not bytes.
317#[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    /// Create an erased input descriptor from typed, initialized storage.
328    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    /// Construct from erased storage.
352    ///
353    /// # Safety
354    /// `data` must be the start of an allocation aligned for `dtype`; the
355    /// allocation must contain `byte_len` initialized, readable bytes for
356    /// `'a`, and every byte extent must represent valid values for `dtype`.
357    /// The allocation and metadata must outlive `'a`; no mutable alias may
358    /// exist while this descriptor is used. Alignment is an allocation
359    /// property and must not be inferred from the observed address alone.
360    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    /// Borrow the initialized storage as its concrete scalar type.
381    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/// Borrowed dtype-erased raw strided output layout.
418///
419/// `dims`, `strides`, and `offset` are expressed in dtype elements, not bytes.
420#[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    /// Create an erased output descriptor from typed, initialized storage.
431    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    /// Construct from erased storage.
456    ///
457    /// # Safety
458    /// `data..data + byte_len` must be an aligned, writable allocation valid
459    /// for `'a`, with valid initialized values for `dtype` throughout the
460    /// complete byte extent; its provenance, extent, lifetime, and exclusive
461    /// aliasing must be upheld by the caller.
462    /// The alignment requirement is on the allocation, not merely the observed
463    /// address. The caller must retain exclusive access: no mutable alias or
464    /// concurrent mutation may exist while this descriptor is used.
465    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    /// Borrow initialized storage as its concrete scalar type.
487    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    /// Borrow initialized storage mutably as its concrete scalar type.
503    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/// Borrowed dtype-erased raw strided output whose reachable elements may be
540/// uninitialized.
541///
542/// This descriptor is only accepted by operations that prove and perform a
543/// full overwrite of every reachable logical destination element. The backing
544/// allocation may contain non-reachable holes, which remain uninitialized.
545#[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    /// Create an uninitialized erased output descriptor from typed storage.
556    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    /// Construct from erased uninitialized storage.
584    ///
585    /// # Safety
586    /// `data..data + byte_len` must be an aligned, writable allocation valid
587    /// for `'a`, with provenance and extent sufficient for all reachable
588    /// elements. The caller must ensure every reachable element is completely
589    /// overwritten before it is read or exposed as initialized storage. The
590    /// allocation's alignment is an allocation property independent of the
591    /// observed address, and the descriptor must have exclusive access with no
592    /// mutable alias or concurrent mutation during use.
593    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    /// Borrow the uninitialized storage as its concrete scalar type.
615    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/// Borrowed raw strided input layout.
662///
663/// Use [`RawStridedRef::new`] for checked construction, or
664/// [`RawStridedRef::new_unchecked`] when a higher-level compiled plan has
665/// already validated bounds.
666#[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    /// Create a raw strided input after validating reachable offsets.
676    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    /// Create a raw strided input without bounds checking.
692    ///
693    /// # Safety
694    /// The caller must ensure every index reachable by `dims`/`strides` from
695    /// `offset` lies inside `data`.
696    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    /// Convert to an immutable owning-metadata view.
740    ///
741    /// This is for compatibility paths. Hot prepared paths should use the raw
742    /// accessors directly and avoid this conversion.
743    #[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/// Borrowed raw strided output layout.
750///
751/// This is the mutable counterpart to [`RawStridedRef`]. It avoids allocating
752/// owned shape/stride metadata in prepared replay paths.
753#[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    /// Create a raw strided output after validating reachable offsets.
763    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    /// Create a raw strided output without bounds checking.
779    ///
780    /// # Safety
781    /// The caller must ensure every index reachable by `dims`/`strides` from
782    /// `offset` lies inside `data`, and no aliases violate mutable access.
783    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    /// Convert to an immutable owning-metadata view.
841    ///
842    /// This is for compatibility paths. Hot prepared paths should use the raw
843    /// accessors directly and avoid this conversion.
844    #[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    /// Convert to a mutable owning-metadata view.
850    ///
851    /// This is for compatibility paths. Hot prepared paths should use the raw
852    /// accessors directly and avoid this conversion.
853    #[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        // Write through the handed off pointer so no reborrow of `bytes`
924        // invalidates it.
925        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}