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}