1use crate::bit_util::ceil;
21
22pub fn set_bits(
34 write_data: &mut [u8],
35 data: &[u8],
36 offset_write: usize,
37 offset_read: usize,
38 len: usize,
39) -> usize {
40 assert!(
41 offset_write
42 .checked_add(len)
43 .expect("operation will overflow write buffer")
44 <= write_data.len() * 8
45 );
46 assert!(
47 offset_read
48 .checked_add(len)
49 .expect("operation will overflow read buffer")
50 <= data.len() * 8
51 );
52 let mut null_count = 0;
53 let mut acc = 0;
54 while len > acc {
55 let (n, len_set) = unsafe {
59 set_upto_64bits(
60 write_data,
61 data,
62 offset_write + acc,
63 offset_read + acc,
64 len - acc,
65 )
66 };
67 null_count += n;
68 acc += len_set;
69 }
70
71 null_count
72}
73
74#[inline]
80unsafe fn set_upto_64bits(
81 write_data: &mut [u8],
82 data: &[u8],
83 offset_write: usize,
84 offset_read: usize,
85 len: usize,
86) -> (usize, usize) {
87 let read_byte = offset_read / 8;
88 let read_shift = offset_read % 8;
89 let write_byte = offset_write / 8;
90 let write_shift = offset_write % 8;
91
92 if len >= 64 {
93 let chunk = unsafe { data.as_ptr().add(read_byte).cast::<u64>().read_unaligned() };
94 if read_shift == 0 {
95 if write_shift == 0 {
96 let len = 64;
98 let null_count = chunk.count_zeros() as usize;
99 unsafe { write_u64_bytes(write_data, write_byte, chunk) };
100 (null_count, len)
101 } else {
102 let len = 64 - write_shift;
104 let chunk = chunk << write_shift;
105 let null_count = len - chunk.count_ones() as usize;
106 unsafe { or_write_u64_bytes(write_data, write_byte, chunk) };
107 (null_count, len)
108 }
109 } else if write_shift == 0 {
110 let len = 64 - 8; let chunk = (chunk >> read_shift) & 0x00FFFFFFFFFFFFFF; let null_count = len - chunk.count_ones() as usize;
114 unsafe { write_u64_bytes(write_data, write_byte, chunk) };
115 (null_count, len)
116 } else {
117 let len = 64 - std::cmp::max(read_shift, write_shift);
118 let chunk = (chunk >> read_shift) << write_shift;
119 let null_count = len - chunk.count_ones() as usize;
120 unsafe { or_write_u64_bytes(write_data, write_byte, chunk) };
121 (null_count, len)
122 }
123 } else if len == 1 {
124 let byte_chunk = (unsafe { data.get_unchecked(read_byte) } >> read_shift) & 1;
125 unsafe { *write_data.get_unchecked_mut(write_byte) |= byte_chunk << write_shift };
126 ((byte_chunk ^ 1) as usize, 1)
127 } else {
128 let len = std::cmp::min(len, 64 - std::cmp::max(read_shift, write_shift));
129 let bytes = ceil(len + read_shift, 8);
130 let chunk = unsafe { read_bytes_to_u64(data, read_byte, bytes) };
132 let mask = u64::MAX >> (64 - len);
133 let chunk = (chunk >> read_shift) & mask; let chunk = chunk << write_shift; let null_count = len - chunk.count_ones() as usize;
136 let bytes = ceil(len + write_shift, 8);
137 for (i, c) in chunk.to_le_bytes().iter().enumerate().take(bytes) {
138 unsafe { *write_data.get_unchecked_mut(write_byte + i) |= c };
139 }
140 (null_count, len)
141 }
142}
143
144#[inline]
147unsafe fn read_bytes_to_u64(data: &[u8], offset: usize, count: usize) -> u64 {
148 debug_assert!(count <= 8);
149 let mut tmp: u64 = 0;
150 let src = unsafe { data.as_ptr().add(offset) };
151 unsafe { std::ptr::copy_nonoverlapping(src, std::ptr::from_mut(&mut tmp).cast::<u8>(), count) };
152 tmp
153}
154
155#[inline]
158unsafe fn write_u64_bytes(data: &mut [u8], offset: usize, chunk: u64) {
159 #[expect(
160 clippy::cast_ptr_alignment,
161 reason = "the pointer is only written through `write_unaligned`"
162 )]
163 let ptr = unsafe { data.as_mut_ptr().add(offset) }.cast::<u64>();
164 unsafe { ptr.write_unaligned(chunk) };
165}
166
167#[inline]
173unsafe fn or_write_u64_bytes(data: &mut [u8], offset: usize, chunk: u64) {
174 let ptr = unsafe { data.as_mut_ptr().add(offset) };
175 let chunk = chunk | (unsafe { *ptr }) as u64;
176 unsafe { ptr.cast::<u64>().write_unaligned(chunk) };
177}
178
179#[cfg(test)]
180mod tests {
181 use super::*;
182 use crate::bit_util::{get_bit, set_bit, unset_bit};
183 use rand::prelude::StdRng;
184 use rand::{RngExt, SeedableRng, TryRng};
185 use std::fmt::Display;
186
187 #[test]
188 fn test_set_bits_aligned() {
189 SetBitsTest {
190 write_data: vec![0, 0, 0, 0, 0, 0, 0, 0, 0, 0],
191 data: vec![
192 0b11100111, 0b10100101, 0b10011001, 0b11011011, 0b11101011, 0b11000011, 0b11100111,
193 0b10100101,
194 ],
195 offset_write: 8,
196 offset_read: 0,
197 len: 64,
198 expected_data: vec![
199 0, 0b11100111, 0b10100101, 0b10011001, 0b11011011, 0b11101011, 0b11000011,
200 0b11100111, 0b10100101, 0,
201 ],
202 expected_null_count: 24,
203 }
204 .verify();
205 }
206
207 #[test]
208 fn test_set_bits_unaligned_destination_start() {
209 SetBitsTest {
210 write_data: vec![0, 0, 0, 0, 0, 0, 0, 0, 0, 0],
211 data: vec![
212 0b11100111, 0b10100101, 0b10011001, 0b11011011, 0b11101011, 0b11000011, 0b11100111,
213 0b10100101,
214 ],
215 offset_write: 3,
216 offset_read: 0,
217 len: 64,
218 expected_data: vec![
219 0b00111000, 0b00101111, 0b11001101, 0b11011100, 0b01011110, 0b00011111, 0b00111110,
220 0b00101111, 0b00000101, 0b00000000,
221 ],
222 expected_null_count: 24,
223 }
224 .verify();
225 }
226
227 #[test]
228 fn test_set_bits_unaligned_destination_end() {
229 SetBitsTest {
230 write_data: vec![0, 0, 0, 0, 0, 0, 0, 0, 0, 0],
231 data: vec![
232 0b11100111, 0b10100101, 0b10011001, 0b11011011, 0b11101011, 0b11000011, 0b11100111,
233 0b10100101,
234 ],
235 offset_write: 8,
236 offset_read: 0,
237 len: 62,
238 expected_data: vec![
239 0, 0b11100111, 0b10100101, 0b10011001, 0b11011011, 0b11101011, 0b11000011,
240 0b11100111, 0b00100101, 0,
241 ],
242 expected_null_count: 23,
243 }
244 .verify();
245 }
246
247 #[test]
248 fn test_set_bits_unaligned() {
249 SetBitsTest {
250 write_data: vec![0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0],
251 data: vec![
252 0b11100111, 0b10100101, 0b10011001, 0b11011011, 0b11101011, 0b11000011, 0b11100111,
253 0b10100101, 0b10011001, 0b11011011, 0b11101011, 0b11000011, 0b11100111, 0b10100101,
254 0b10011001, 0b11011011, 0b11101011, 0b11000011,
255 ],
256 offset_write: 3,
257 offset_read: 5,
258 len: 95,
259 expected_data: vec![
260 0b01111000, 0b01101001, 0b11100110, 0b11110110, 0b11111010, 0b11110000, 0b01111001,
261 0b01101001, 0b11100110, 0b11110110, 0b11111010, 0b11110000, 0b00000001,
262 ],
263 expected_null_count: 35,
264 }
265 .verify();
266 }
267
268 #[test]
269 fn set_bits_fuzz() {
270 let mut rng = StdRng::seed_from_u64(42);
271 let mut data = SetBitsTest::new();
272 for _ in 0..100 {
273 data.regen(&mut rng);
274 data.verify();
275 }
276 }
277
278 #[derive(Debug, Default)]
279 struct SetBitsTest {
280 write_data: Vec<u8>,
282 data: Vec<u8>,
284 offset_write: usize,
285 offset_read: usize,
286 len: usize,
287 expected_data: Vec<u8>,
289 expected_null_count: usize,
291 }
292
293 struct BinaryFormatter<'a>(&'a [u8]);
295 impl Display for BinaryFormatter<'_> {
296 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
297 for byte in self.0 {
298 write!(f, "{byte:08b} ")?;
299 }
300 write!(f, " ")?;
301 Ok(())
302 }
303 }
304
305 impl Display for SetBitsTest {
306 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
307 writeln!(f, "SetBitsTest {{")?;
308 writeln!(f, " write_data: {}", BinaryFormatter(&self.write_data))?;
309 writeln!(f, " data: {}", BinaryFormatter(&self.data))?;
310 writeln!(
311 f,
312 " expected_data: {}",
313 BinaryFormatter(&self.expected_data)
314 )?;
315 writeln!(f, " offset_write: {}", self.offset_write)?;
316 writeln!(f, " offset_read: {}", self.offset_read)?;
317 writeln!(f, " len: {}", self.len)?;
318 writeln!(f, " expected_null_count: {}", self.expected_null_count)?;
319 writeln!(f, "}}")
320 }
321 }
322
323 impl SetBitsTest {
324 fn new() -> Self {
326 Self::default()
327 }
328
329 fn regen(&mut self, rng: &mut StdRng) {
331 let len = rng.random_range(0..=200);
343
344 let offset_write_bits = rng.random_range(0..=200);
346 let offset_write_bytes = if offset_write_bits % 8 == 0 {
347 offset_write_bits / 8
348 } else {
349 (offset_write_bits / 8) + 1
350 };
351 let extra_write_data_bytes = rng.random_range(0..=5); let extra_read_data_bytes = rng.random_range(0..=5); let offset_read_bits = rng.random_range(0..=200);
356 let offset_read_bytes = if offset_read_bits % 8 != 0 {
357 (offset_read_bits / 8) + 1
358 } else {
359 offset_read_bits / 8
360 };
361
362 self.write_data.clear();
364 self.write_data
365 .resize(offset_write_bytes + len + extra_write_data_bytes, 0);
366
367 self.offset_write = offset_write_bits;
371
372 self.data
374 .resize(offset_read_bytes + len + extra_read_data_bytes, 0);
375 rng.try_fill_bytes(self.data.as_mut_slice()).unwrap();
377 self.offset_read = offset_read_bits;
378
379 self.len = len;
380
381 self.expected_data.resize(self.write_data.len(), 0);
383 self.expected_data.copy_from_slice(&self.write_data);
384
385 self.expected_null_count = 0;
386 for i in 0..self.len {
387 let bit = get_bit(&self.data, self.offset_read + i);
388 if bit {
389 set_bit(&mut self.expected_data, self.offset_write + i);
390 } else {
391 unset_bit(&mut self.expected_data, self.offset_write + i);
392 self.expected_null_count += 1;
393 }
394 }
395 }
396
397 fn verify(&self) {
399 let mut actual = self.write_data.clone();
401 let null_count = set_bits(
402 &mut actual,
403 &self.data,
404 self.offset_write,
405 self.offset_read,
406 self.len,
407 );
408
409 assert_eq!(actual, self.expected_data, "self: {self}");
410 assert_eq!(null_count, self.expected_null_count, "self: {self}");
411 }
412 }
413
414 #[test]
415 fn test_set_upto_64bits() {
416 let write_data: &mut [u8] = &mut [0; 9];
418 let data: &[u8] = &[
419 0b00000001, 0b00000001, 0b00000001, 0b00000001, 0b00000001, 0b00000001, 0b00000001,
420 0b00000001, 0b00000001,
421 ];
422 let offset_write = 1;
423 let offset_read = 0;
424 let len = 65;
425 let (n, len_set) =
426 unsafe { set_upto_64bits(write_data, data, offset_write, offset_read, len) };
427 assert_eq!(n, 55);
428 assert_eq!(len_set, 63);
429 assert_eq!(
430 write_data,
431 &[
432 0b00000010, 0b00000010, 0b00000010, 0b00000010, 0b00000010, 0b00000010, 0b00000010,
433 0b00000010, 0b00000000
434 ]
435 );
436
437 let write_data: &mut [u8] = &mut [0b00000000];
439 let data: &[u8] = &[0b00000001];
440 let offset_write = 1;
441 let offset_read = 0;
442 let len = 1;
443 let (n, len_set) =
444 unsafe { set_upto_64bits(write_data, data, offset_write, offset_read, len) };
445 assert_eq!(n, 0);
446 assert_eq!(len_set, 1);
447 assert_eq!(write_data, &[0b00000010]);
448 }
449
450 #[test]
451 #[should_panic(expected = "operation will overflow read buffer")]
452 fn test_overflow_read_buffer_bounds() {
453 let data = [0u8; 1];
455 let mut write_data = [0u8; 1];
456
457 let offset_write: usize = 0;
461 let offset_read: usize = usize::MAX - 7;
462 let len: usize = 8;
463
464 let _nulls = set_bits(&mut write_data, &data, offset_write, offset_read, len);
466 }
467
468 #[test]
469 #[should_panic(expected = "operation will overflow write buffer")]
470 fn test_overflow_write_buffer_bounds() {
471 let data = [0u8; 1];
473 let mut write_data = [0u8; 1];
474
475 let offset_write: usize = usize::MAX - 7;
479 let offset_read: usize = 0;
480 let len: usize = 8;
481
482 let _nulls = set_bits(&mut write_data, &data, offset_write, offset_read, len);
484 }
485}