Skip to main content

diskann_vector/
unaligned.rs

1/*
2 * Copyright (c) Microsoft Corporation.
3 * Licensed under the MIT license.
4 */
5
6use std::marker::PhantomData;
7
8/// A minimally functional span over a potentially unaligned slice of elements of type `T`.
9///
10/// Like `&[T]`, this type guarantees that the memory
11/// `[self.as_ptr(), self.as_ptr().add(self.len()))` is valid for reads.
12///
13/// However, unlike `&[T]`, the pointer [`Self::as_ptr`] is **not** guaranteed to be aligned
14/// to `std::mem::align_of::<T>()`.
15///
16/// If the type `T` is [`Copy`], then [`std::ptr::read_unaligned`] can be used on the valid
17/// memory region for this slice to access values of type `T`.
18#[derive(Debug)]
19pub struct UnalignedSlice<'a, T> {
20    ptr: *const T,
21    len: usize,
22    _lifetime: PhantomData<&'a T>,
23}
24
25/// SAFETY: `UnalignedSlice` is like a `&[T]`: it can be sent to other threads when `&T: Send`,
26/// which holds iff `T: Sync`.
27unsafe impl<T> Send for UnalignedSlice<'_, T> where T: Sync {}
28
29/// SAFETY: `UnalignedSlice` is like a `&[T]`: it can be shared with other threads when `&T`
30/// is send, implying the bound `T: Sync`.
31unsafe impl<T> Sync for UnalignedSlice<'_, T> where T: Sync {}
32
33impl<'a, T> UnalignedSlice<'a, T> {
34    /// Construct a new [`UnalignedSlice`] over the region `[ptr, ptr.add(len))`.
35    ///
36    /// # Safety
37    ///
38    /// The provided memory region must be valid for reading via [`std::ptr::read_unaligned`]
39    /// for the duration of the associated lifetime.
40    ///
41    /// Note that `ptr` need not be aligned to `std::mem::align_of::<T>()`.
42    ///
43    /// Argument `ptr` may be null if `len == 0`.
44    pub const unsafe fn new(ptr: *const T, len: usize) -> Self {
45        Self {
46            ptr,
47            len,
48            _lifetime: PhantomData,
49        }
50    }
51
52    /// Return the number of elements of type `T` available for reading from [`Self::as_ptr`].
53    pub const fn len(&self) -> usize {
54        self.len
55    }
56
57    /// Return `true` only if the associated slice is empty.
58    pub const fn is_empty(&self) -> bool {
59        self.len() == 0
60    }
61
62    /// Return the base pointer of the slice.
63    ///
64    /// **NOTE**: It is **not** guaranteed that the returned pointer is aligned!
65    pub const fn as_ptr(&self) -> *const T {
66        self.ptr
67    }
68}
69
70impl<T> Clone for UnalignedSlice<'_, T> {
71    fn clone(&self) -> Self {
72        *self
73    }
74}
75
76impl<T> Copy for UnalignedSlice<'_, T> {}
77
78impl<'a, T> From<&'a [T]> for UnalignedSlice<'a, T> {
79    fn from(slice: &'a [T]) -> Self {
80        // SAFETY: Slices are inherently valid, so this construction is safe.
81        unsafe { Self::new(slice.as_ptr(), slice.len()) }
82    }
83}
84
85impl<'a, T, const N: usize> From<&'a [T; N]> for UnalignedSlice<'a, T> {
86    fn from(slice: &'a [T; N]) -> Self {
87        // SAFETY: Slices are inherently valid, so this construction is safe.
88        unsafe { Self::new(slice.as_ptr(), N) }
89    }
90}
91
92/// View `self` as an [`UnalignedSlice`].
93pub trait AsUnaligned {
94    /// The element type of the slice.
95    type Element;
96
97    /// Return an [`UnalignedSlice`] view of `self`.
98    fn as_unaligned(&self) -> UnalignedSlice<'_, Self::Element>;
99}
100
101impl<T> AsUnaligned for UnalignedSlice<'_, T> {
102    type Element = T;
103    fn as_unaligned(&self) -> UnalignedSlice<'_, T> {
104        *self
105    }
106}
107
108impl<T> AsUnaligned for &[T] {
109    type Element = T;
110    fn as_unaligned(&self) -> UnalignedSlice<'_, T> {
111        (*self).into()
112    }
113}
114
115impl<T, const N: usize> AsUnaligned for &[T; N] {
116    type Element = T;
117    fn as_unaligned(&self) -> UnalignedSlice<'_, T> {
118        (*self).into()
119    }
120}
121
122impl<T, const N: usize> AsUnaligned for [T; N] {
123    type Element = T;
124    fn as_unaligned(&self) -> UnalignedSlice<'_, T> {
125        self.into()
126    }
127}
128
129/// A utility that offsets a collection of `T` by one byte to test the guarantee that
130/// distance functions can operate on unaligned pointers.
131///
132/// # Invariant
133///
134/// We maintain the following invariants
135/// * `self.data.as_ptr().add(1)` is always safe (i.e., `data.len() >= 1`)
136/// * `data.len() == self.len * std::mem::size_of::<T>() + 1`: We can construct an
137///   [`UnalignedSlice`] of length `self.len` starting from `self.data.as_ptr().add(1)`.
138#[cfg(test)]
139#[derive(Debug)]
140pub(crate) struct Buffer<T>
141where
142    T: bytemuck::Pod,
143{
144    data: Vec<u8>,
145    len: usize,
146    _type: PhantomData<T>,
147}
148
149#[cfg(test)]
150impl<T> Default for Buffer<T>
151where
152    T: bytemuck::Pod,
153{
154    fn default() -> Self {
155        Self {
156            // NOTE: Length 1 is important to maintain struct invariants.
157            data: vec![0u8; 1],
158            len: 0,
159            _type: PhantomData,
160        }
161    }
162}
163
164#[cfg(test)]
165impl<T> Buffer<T>
166where
167    T: bytemuck::Pod,
168{
169    pub(crate) fn new(x: &[T]) -> Self {
170        let mut this = Self::default();
171        this.copy(x);
172        this
173    }
174
175    pub(crate) fn copy(&mut self, x: &[T]) {
176        let bytes = std::mem::size_of_val(x);
177        self.data.resize(bytes.checked_add(1).unwrap(), 0u8);
178
179        // SAFETY: We maintain the invariant that `data.len() >= 1`, so `ptr::add` is valid.
180        let dst = unsafe { self.data.as_mut_ptr().add(1) };
181
182        // SAFETY: The `bytemuck::Pod` bound guarantees `Copy`, so we need not worry about
183        // destructors. Further:
184        //
185        // * `src` is valid for `bytes` since that is how we obtained the value `bytes` in
186        //   the first place.
187        // * `dst` is valid for `bytes` since we just resized.
188        // * Both `src` and `dst` are `u8` and are trivially aligned.
189        // * Neither region can overlap because we borrow self by mutable reference -
190        //   therefore `data` must be disjoint from the memory for `x`.
191        unsafe {
192            std::ptr::copy_nonoverlapping::<u8>(
193                bytemuck::must_cast_slice::<T, u8>(x).as_ptr(),
194                dst,
195                bytes,
196            );
197        }
198
199        self.len = x.len();
200    }
201
202    pub(crate) fn as_unaligned(&self) -> UnalignedSlice<'_, T> {
203        // SAFETY: The invariants maintained by this class guarantee the validity of the
204        // memory of the returned slice.
205        //
206        // The `bytemuck::Pod` bound means the resulting slice is useable via
207        // [`ptr::read_unaligned`].
208        unsafe { UnalignedSlice::new(self.data.as_ptr().add(1).cast::<T>(), self.len) }
209    }
210}