diff --git a/vortex-buffer/benches/vortex_buffer.rs b/vortex-buffer/benches/vortex_buffer.rs index 3ba571751f6..81bae392b14 100644 --- a/vortex-buffer/benches/vortex_buffer.rs +++ b/vortex-buffer/benches/vortex_buffer.rs @@ -75,7 +75,7 @@ impl MapEach for Arrow MapEach for Buffer { +impl MapEach for Buffer { type Output = BufferMut; fn map_each(self, f: F) -> Self::Output @@ -86,7 +86,7 @@ impl MapEach for Buffer { } } -impl MapEach for BufferMut { +impl MapEach for BufferMut { type Output = BufferMut; fn map_each(self, f: F) -> Self::Output diff --git a/vortex-buffer/src/buffer.rs b/vortex-buffer/src/buffer.rs index 9e954bdcde4..ab0ca530f00 100644 --- a/vortex-buffer/src/buffer.rs +++ b/vortex-buffer/src/buffer.rs @@ -345,6 +345,7 @@ impl Buffer { pub fn map_each_in_place(self, mut f: F) -> BufferMut where T: Copy, + R: Copy, F: FnMut(T) -> R, { match self.try_into_mut() { diff --git a/vortex-buffer/src/buffer_mut.rs b/vortex-buffer/src/buffer_mut.rs index 136267b11b8..0b45fe23772 100644 --- a/vortex-buffer/src/buffer_mut.rs +++ b/vortex-buffer/src/buffer_mut.rs @@ -785,6 +785,7 @@ impl BufferMut { pub fn map_each_in_place(self, mut f: F) -> BufferMut where T: Copy, + R: Copy, F: FnMut(T) -> R, { assert_eq!( @@ -796,7 +797,9 @@ impl BufferMut { let mut buf: BufferMut = unsafe { std::mem::transmute(self) }; buf.iter_mut() .for_each(|item| *item = f(unsafe { std::mem::transmute_copy(item) })); - buf + // `transmute` preserves `T`'s alignment, which can be weaker than `R`'s. + let alignment = buf.alignment().max(Alignment::of::()); + buf.aligned(alignment) } /// Return a `BufferMut` with the same data as this one with the given alignment. @@ -1074,6 +1077,7 @@ fn misaligned_scalar_type(alignment: Alignment, scalar_align: Alignment) -> ! { #[cfg(test)] mod tests { + use std::alloc::Layout; use std::cell::Cell; use allocator_api2::alloc::Global; @@ -1298,6 +1302,16 @@ mod tests { assert_eq!(buf.as_slice(), &[1u32, 2, 3]); } + #[test] + fn map_each_in_place_keeps_output_alignment_valid() { + let bytes = BufferMut::<[u8; 4]>::copy_from_aligned([[1u8, 0, 0, 0]], Alignment::new(1)); + assert_eq!(1, Layout::new::<[u8; 4]>().align()); + let words = bytes.map_each_in_place(u32::from_ne_bytes); + + assert_eq!(words.as_slice(), [1]); + assert!(words.alignment().is_aligned_to(Alignment::of::())); + } + #[test] fn buffer_mut_zeroed() { const LEN: usize = 17;