1#![allow(unsafe_op_in_unsafe_fn)]
4
5use crate::gf;
17
18pub struct GfMulTables {
20 pub lo_lo: [u8; 16],
21 pub lo_hi: [u8; 16],
22 pub hi_lo: [u8; 16],
23 pub hi_hi: [u8; 16],
24 pub ulo_lo: [u8; 16],
25 pub ulo_hi: [u8; 16],
26 pub uhi_lo: [u8; 16],
27 pub uhi_hi: [u8; 16],
28}
29
30impl GfMulTables {
31 pub fn new(constant: u16) -> Self {
32 let mut t = GfMulTables {
33 lo_lo: [0; 16],
34 lo_hi: [0; 16],
35 hi_lo: [0; 16],
36 hi_hi: [0; 16],
37 ulo_lo: [0; 16],
38 ulo_hi: [0; 16],
39 uhi_lo: [0; 16],
40 uhi_hi: [0; 16],
41 };
42 for i in 0..16u16 {
43 let v = gf::mul(constant, i);
44 t.lo_lo[i as usize] = v as u8;
45 t.lo_hi[i as usize] = (v >> 8) as u8;
46 let v = gf::mul(constant, i << 4);
47 t.hi_lo[i as usize] = v as u8;
48 t.hi_hi[i as usize] = (v >> 8) as u8;
49 let v = gf::mul(constant, i << 8);
50 t.ulo_lo[i as usize] = v as u8;
51 t.ulo_hi[i as usize] = (v >> 8) as u8;
52 let v = gf::mul(constant, i << 12);
53 t.uhi_lo[i as usize] = v as u8;
54 t.uhi_hi[i as usize] = (v >> 8) as u8;
55 }
56 t
57 }
58}
59
60pub fn mul_add_buffer(dst: &mut [u8], src: &[u8], constant: u16) {
66 assert_eq!(dst.len(), src.len());
67 if constant == 0 {
68 return;
69 }
70 if constant == 1 {
71 xor_buffers(dst, src);
72 return;
73 }
74
75 #[cfg(target_arch = "x86_64")]
76 {
77 if is_x86_feature_detected!("avx2") {
78 unsafe { mul_add_buffer_avx2(dst, src, constant) };
79 return;
80 }
81 if is_x86_feature_detected!("ssse3") {
82 unsafe { mul_add_buffer_ssse3(dst, src, constant) };
83 return;
84 }
85 }
86 mul_add_buffer_scalar(dst, src, constant);
87}
88
89pub fn mul_add_multi(dst: &mut [u8], srcs: &[&[u8]], coeffs: &[u16]) {
101 assert_eq!(srcs.len(), coeffs.len());
102
103 let active: Vec<(usize, u16)> = coeffs
105 .iter()
106 .copied()
107 .enumerate()
108 .filter(|(_, c)| *c != 0)
109 .collect();
110
111 if active.is_empty() {
112 return;
113 }
114
115 #[cfg(target_arch = "x86_64")]
116 {
117 if is_x86_feature_detected!("avx2") {
118 let mut i = 0;
120 while i + 1 < active.len() {
121 let (idx1, c1) = active[i];
122 let (idx2, c2) = active[i + 1];
123 unsafe { mul_add_pair_avx2(dst, srcs[idx1], c1, srcs[idx2], c2) };
124 i += 2;
125 }
126 if i < active.len() {
128 let (idx, c) = active[i];
129 unsafe { mul_add_buffer_avx2(dst, srcs[idx], c) };
130 }
131 return;
132 }
133 }
134
135 for &(idx, coeff) in &active {
136 mul_add_buffer(dst, srcs[idx], coeff);
137 }
138}
139
140pub fn xor_buffers(dst: &mut [u8], src: &[u8]) {
142 assert_eq!(dst.len(), src.len());
143 #[cfg(target_arch = "x86_64")]
144 {
145 if is_x86_feature_detected!("avx2") {
146 unsafe { xor_buffers_avx2(dst, src) };
147 return;
148 }
149 }
150 for (d, s) in dst.iter_mut().zip(src.iter()) {
151 *d ^= s;
152 }
153}
154
155fn mul_add_buffer_scalar(dst: &mut [u8], src: &[u8], constant: u16) {
160 let len = dst.len() / 2;
161 for i in 0..len {
162 let off = i * 2;
163 let s = u16::from_le_bytes([src[off], src[off + 1]]);
164 let d = u16::from_le_bytes([dst[off], dst[off + 1]]);
165 let result = d ^ gf::mul(constant, s);
166 dst[off] = result as u8;
167 dst[off + 1] = (result >> 8) as u8;
168 }
169}
170
171#[cfg(target_arch = "x86_64")]
176#[target_feature(enable = "avx2")]
177unsafe fn mul_add_buffer_avx2(dst: &mut [u8], src: &[u8], constant: u16) {
178 let tables = GfMulTables::new(constant);
179 gf_mul_add_avx2_inner(dst, src, &tables);
180}
181
182#[cfg(target_arch = "x86_64")]
184#[target_feature(enable = "avx2")]
185unsafe fn gf_mul_add_avx2_inner(dst: &mut [u8], src: &[u8], tables: &GfMulTables) {
186 use std::arch::x86_64::*;
187
188 let nibble_mask = _mm256_set1_epi8(0x0F);
190 let tbl_lo_lo = _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.lo_lo.as_ptr() as *const _));
191 let tbl_lo_hi = _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.lo_hi.as_ptr() as *const _));
192 let tbl_hi_lo = _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.hi_lo.as_ptr() as *const _));
193 let tbl_hi_hi = _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.hi_hi.as_ptr() as *const _));
194 let tbl_ulo_lo =
195 _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.ulo_lo.as_ptr() as *const _));
196 let tbl_ulo_hi =
197 _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.ulo_hi.as_ptr() as *const _));
198 let tbl_uhi_lo =
199 _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.uhi_lo.as_ptr() as *const _));
200 let tbl_uhi_hi =
201 _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.uhi_hi.as_ptr() as *const _));
202
203 let deint_lo = _mm256_broadcastsi128_si256(_mm_setr_epi8(
206 0, 2, 4, 6, 8, 10, 12, 14, -1, -1, -1, -1, -1, -1, -1, -1,
207 ));
208 let deint_hi = _mm256_broadcastsi128_si256(_mm_setr_epi8(
209 1, 3, 5, 7, 9, 11, 13, 15, -1, -1, -1, -1, -1, -1, -1, -1,
210 ));
211
212 let len = dst.len();
213 let chunks = len / 32;
214
215 for chunk in 0..chunks {
216 let off = chunk * 32;
217 let src_data = _mm256_loadu_si256(src[off..].as_ptr() as *const __m256i);
218
219 let src_lo_bytes = _mm256_shuffle_epi8(src_data, deint_lo);
221 let src_hi_bytes = _mm256_shuffle_epi8(src_data, deint_hi);
222
223 let lo_nib = _mm256_and_si256(src_lo_bytes, nibble_mask);
225 let hi_nib = _mm256_and_si256(_mm256_srli_epi16(src_lo_bytes, 4), nibble_mask);
226 let ulo_nib = _mm256_and_si256(src_hi_bytes, nibble_mask);
227 let uhi_nib = _mm256_and_si256(_mm256_srli_epi16(src_hi_bytes, 4), nibble_mask);
228
229 let r_lo = _mm256_xor_si256(
231 _mm256_xor_si256(
232 _mm256_shuffle_epi8(tbl_lo_lo, lo_nib),
233 _mm256_shuffle_epi8(tbl_hi_lo, hi_nib),
234 ),
235 _mm256_xor_si256(
236 _mm256_shuffle_epi8(tbl_ulo_lo, ulo_nib),
237 _mm256_shuffle_epi8(tbl_uhi_lo, uhi_nib),
238 ),
239 );
240 let r_hi = _mm256_xor_si256(
242 _mm256_xor_si256(
243 _mm256_shuffle_epi8(tbl_lo_hi, lo_nib),
244 _mm256_shuffle_epi8(tbl_hi_hi, hi_nib),
245 ),
246 _mm256_xor_si256(
247 _mm256_shuffle_epi8(tbl_ulo_hi, ulo_nib),
248 _mm256_shuffle_epi8(tbl_uhi_hi, uhi_nib),
249 ),
250 );
251
252 let result = _mm256_unpacklo_epi8(r_lo, r_hi);
254
255 let dst_val = _mm256_loadu_si256(dst[off..].as_ptr() as *const __m256i);
257 _mm256_storeu_si256(
258 dst[off..].as_mut_ptr() as *mut __m256i,
259 _mm256_xor_si256(dst_val, result),
260 );
261 }
262
263 let rem = chunks * 32;
265 if rem < len {
266 mul_add_buffer_scalar(
267 &mut dst[rem..],
268 &src[rem..],
269 gf::mul(
270 tables.lo_lo[1] as u16 | ((tables.lo_hi[1] as u16) << 8),
274 1,
275 ),
276 );
277 }
279}
280
281#[cfg(target_arch = "x86_64")]
283#[target_feature(enable = "avx2")]
284unsafe fn mul_add_pair_avx2(dst: &mut [u8], src1: &[u8], c1: u16, src2: &[u8], c2: u16) {
285 use std::arch::x86_64::*;
286
287 let t1 = GfMulTables::new(c1);
288 let t2 = GfMulTables::new(c2);
289
290 let nibble_mask = _mm256_set1_epi8(0x0F);
291 let deint_lo = _mm256_broadcastsi128_si256(_mm_setr_epi8(
292 0, 2, 4, 6, 8, 10, 12, 14, -1, -1, -1, -1, -1, -1, -1, -1,
293 ));
294 let deint_hi = _mm256_broadcastsi128_si256(_mm_setr_epi8(
295 1, 3, 5, 7, 9, 11, 13, 15, -1, -1, -1, -1, -1, -1, -1, -1,
296 ));
297
298 let t1_lo_lo = _mm256_broadcastsi128_si256(_mm_loadu_si128(t1.lo_lo.as_ptr() as *const _));
300 let t1_lo_hi = _mm256_broadcastsi128_si256(_mm_loadu_si128(t1.lo_hi.as_ptr() as *const _));
301 let t1_hi_lo = _mm256_broadcastsi128_si256(_mm_loadu_si128(t1.hi_lo.as_ptr() as *const _));
302 let t1_hi_hi = _mm256_broadcastsi128_si256(_mm_loadu_si128(t1.hi_hi.as_ptr() as *const _));
303 let t1_ulo_lo = _mm256_broadcastsi128_si256(_mm_loadu_si128(t1.ulo_lo.as_ptr() as *const _));
304 let t1_ulo_hi = _mm256_broadcastsi128_si256(_mm_loadu_si128(t1.ulo_hi.as_ptr() as *const _));
305 let t1_uhi_lo = _mm256_broadcastsi128_si256(_mm_loadu_si128(t1.uhi_lo.as_ptr() as *const _));
306 let t1_uhi_hi = _mm256_broadcastsi128_si256(_mm_loadu_si128(t1.uhi_hi.as_ptr() as *const _));
307
308 let t2_lo_lo = _mm256_broadcastsi128_si256(_mm_loadu_si128(t2.lo_lo.as_ptr() as *const _));
310 let t2_lo_hi = _mm256_broadcastsi128_si256(_mm_loadu_si128(t2.lo_hi.as_ptr() as *const _));
311 let t2_hi_lo = _mm256_broadcastsi128_si256(_mm_loadu_si128(t2.hi_lo.as_ptr() as *const _));
312 let t2_hi_hi = _mm256_broadcastsi128_si256(_mm_loadu_si128(t2.hi_hi.as_ptr() as *const _));
313 let t2_ulo_lo = _mm256_broadcastsi128_si256(_mm_loadu_si128(t2.ulo_lo.as_ptr() as *const _));
314 let t2_ulo_hi = _mm256_broadcastsi128_si256(_mm_loadu_si128(t2.ulo_hi.as_ptr() as *const _));
315 let t2_uhi_lo = _mm256_broadcastsi128_si256(_mm_loadu_si128(t2.uhi_lo.as_ptr() as *const _));
316 let t2_uhi_hi = _mm256_broadcastsi128_si256(_mm_loadu_si128(t2.uhi_hi.as_ptr() as *const _));
317
318 let len = dst.len();
319 let chunks = len / 32;
320
321 for chunk in 0..chunks {
322 let off = chunk * 32;
323
324 let mut acc = _mm256_loadu_si256(dst[off..].as_ptr() as *const __m256i);
326
327 let s1 = _mm256_loadu_si256(src1[off..].as_ptr() as *const __m256i);
329 let s1_lo = _mm256_shuffle_epi8(s1, deint_lo);
330 let s1_hi = _mm256_shuffle_epi8(s1, deint_hi);
331 let n1 = _mm256_and_si256(s1_lo, nibble_mask);
332 let n2 = _mm256_and_si256(_mm256_srli_epi16(s1_lo, 4), nibble_mask);
333 let n3 = _mm256_and_si256(s1_hi, nibble_mask);
334 let n4 = _mm256_and_si256(_mm256_srli_epi16(s1_hi, 4), nibble_mask);
335 let r1_lo = _mm256_xor_si256(
336 _mm256_xor_si256(
337 _mm256_shuffle_epi8(t1_lo_lo, n1),
338 _mm256_shuffle_epi8(t1_hi_lo, n2),
339 ),
340 _mm256_xor_si256(
341 _mm256_shuffle_epi8(t1_ulo_lo, n3),
342 _mm256_shuffle_epi8(t1_uhi_lo, n4),
343 ),
344 );
345 let r1_hi = _mm256_xor_si256(
346 _mm256_xor_si256(
347 _mm256_shuffle_epi8(t1_lo_hi, n1),
348 _mm256_shuffle_epi8(t1_hi_hi, n2),
349 ),
350 _mm256_xor_si256(
351 _mm256_shuffle_epi8(t1_ulo_hi, n3),
352 _mm256_shuffle_epi8(t1_uhi_hi, n4),
353 ),
354 );
355 acc = _mm256_xor_si256(acc, _mm256_unpacklo_epi8(r1_lo, r1_hi));
356
357 let s2 = _mm256_loadu_si256(src2[off..].as_ptr() as *const __m256i);
359 let s2_lo = _mm256_shuffle_epi8(s2, deint_lo);
360 let s2_hi = _mm256_shuffle_epi8(s2, deint_hi);
361 let n1 = _mm256_and_si256(s2_lo, nibble_mask);
362 let n2 = _mm256_and_si256(_mm256_srli_epi16(s2_lo, 4), nibble_mask);
363 let n3 = _mm256_and_si256(s2_hi, nibble_mask);
364 let n4 = _mm256_and_si256(_mm256_srli_epi16(s2_hi, 4), nibble_mask);
365 let r2_lo = _mm256_xor_si256(
366 _mm256_xor_si256(
367 _mm256_shuffle_epi8(t2_lo_lo, n1),
368 _mm256_shuffle_epi8(t2_hi_lo, n2),
369 ),
370 _mm256_xor_si256(
371 _mm256_shuffle_epi8(t2_ulo_lo, n3),
372 _mm256_shuffle_epi8(t2_uhi_lo, n4),
373 ),
374 );
375 let r2_hi = _mm256_xor_si256(
376 _mm256_xor_si256(
377 _mm256_shuffle_epi8(t2_lo_hi, n1),
378 _mm256_shuffle_epi8(t2_hi_hi, n2),
379 ),
380 _mm256_xor_si256(
381 _mm256_shuffle_epi8(t2_ulo_hi, n3),
382 _mm256_shuffle_epi8(t2_uhi_hi, n4),
383 ),
384 );
385 acc = _mm256_xor_si256(acc, _mm256_unpacklo_epi8(r2_lo, r2_hi));
386
387 _mm256_storeu_si256(dst[off..].as_mut_ptr() as *mut __m256i, acc);
389 }
390
391 let rem = chunks * 32;
392 if rem < len {
393 mul_add_buffer_scalar(&mut dst[rem..], &src1[rem..], c1);
394 mul_add_buffer_scalar(&mut dst[rem..], &src2[rem..], c2);
395 }
396}
397
398#[cfg(target_arch = "x86_64")]
400#[target_feature(enable = "avx2")]
401#[allow(dead_code)]
402unsafe fn mul_add_multi_avx2(dst: &mut [u8], srcs: &[&[u8]], active: &[(usize, u16)]) {
403 use std::arch::x86_64::*;
404
405 let nibble_mask = _mm256_set1_epi8(0x0F);
406 let deint_lo = _mm256_broadcastsi128_si256(_mm_setr_epi8(
407 0, 2, 4, 6, 8, 10, 12, 14, -1, -1, -1, -1, -1, -1, -1, -1,
408 ));
409 let deint_hi = _mm256_broadcastsi128_si256(_mm_setr_epi8(
410 1, 3, 5, 7, 9, 11, 13, 15, -1, -1, -1, -1, -1, -1, -1, -1,
411 ));
412
413 let all_tables: Vec<GfMulTables> = active.iter().map(|&(_, c)| GfMulTables::new(c)).collect();
415
416 let len = dst.len();
417 let chunks = len / 32;
418
419 for chunk in 0..chunks {
420 let off = chunk * 32;
421
422 let mut acc = _mm256_loadu_si256(dst[off..].as_ptr() as *const __m256i);
424
425 for (src_i, &(src_idx, _)) in active.iter().enumerate() {
427 let tables = &all_tables[src_i];
428 let src = srcs[src_idx];
429
430 let tbl_lo_lo =
431 _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.lo_lo.as_ptr() as *const _));
432 let tbl_lo_hi =
433 _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.lo_hi.as_ptr() as *const _));
434 let tbl_hi_lo =
435 _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.hi_lo.as_ptr() as *const _));
436 let tbl_hi_hi =
437 _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.hi_hi.as_ptr() as *const _));
438 let tbl_ulo_lo =
439 _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.ulo_lo.as_ptr() as *const _));
440 let tbl_ulo_hi =
441 _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.ulo_hi.as_ptr() as *const _));
442 let tbl_uhi_lo =
443 _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.uhi_lo.as_ptr() as *const _));
444 let tbl_uhi_hi =
445 _mm256_broadcastsi128_si256(_mm_loadu_si128(tables.uhi_hi.as_ptr() as *const _));
446
447 let src_data = _mm256_loadu_si256(src[off..].as_ptr() as *const __m256i);
448
449 let src_lo_bytes = _mm256_shuffle_epi8(src_data, deint_lo);
450 let src_hi_bytes = _mm256_shuffle_epi8(src_data, deint_hi);
451
452 let lo_nib = _mm256_and_si256(src_lo_bytes, nibble_mask);
453 let hi_nib = _mm256_and_si256(_mm256_srli_epi16(src_lo_bytes, 4), nibble_mask);
454 let ulo_nib = _mm256_and_si256(src_hi_bytes, nibble_mask);
455 let uhi_nib = _mm256_and_si256(_mm256_srli_epi16(src_hi_bytes, 4), nibble_mask);
456
457 let r_lo = _mm256_xor_si256(
458 _mm256_xor_si256(
459 _mm256_shuffle_epi8(tbl_lo_lo, lo_nib),
460 _mm256_shuffle_epi8(tbl_hi_lo, hi_nib),
461 ),
462 _mm256_xor_si256(
463 _mm256_shuffle_epi8(tbl_ulo_lo, ulo_nib),
464 _mm256_shuffle_epi8(tbl_uhi_lo, uhi_nib),
465 ),
466 );
467 let r_hi = _mm256_xor_si256(
468 _mm256_xor_si256(
469 _mm256_shuffle_epi8(tbl_lo_hi, lo_nib),
470 _mm256_shuffle_epi8(tbl_hi_hi, hi_nib),
471 ),
472 _mm256_xor_si256(
473 _mm256_shuffle_epi8(tbl_ulo_hi, ulo_nib),
474 _mm256_shuffle_epi8(tbl_uhi_hi, uhi_nib),
475 ),
476 );
477
478 let result = _mm256_unpacklo_epi8(r_lo, r_hi);
479 acc = _mm256_xor_si256(acc, result);
480 }
481
482 _mm256_storeu_si256(dst[off..].as_mut_ptr() as *mut __m256i, acc);
484 }
485
486 let rem = chunks * 32;
488 if rem < len {
489 for &(src_idx, coeff) in active {
490 mul_add_buffer_scalar(&mut dst[rem..], &srcs[src_idx][rem..], coeff);
491 }
492 }
493}
494
495#[cfg(target_arch = "x86_64")]
496#[target_feature(enable = "avx2")]
497unsafe fn xor_buffers_avx2(dst: &mut [u8], src: &[u8]) {
498 use std::arch::x86_64::*;
499 let len = dst.len();
500 let chunks = len / 32;
501 for chunk in 0..chunks {
502 let off = chunk * 32;
503 let s = _mm256_loadu_si256(src[off..].as_ptr() as *const __m256i);
504 let d = _mm256_loadu_si256(dst[off..].as_ptr() as *const __m256i);
505 _mm256_storeu_si256(
506 dst[off..].as_mut_ptr() as *mut __m256i,
507 _mm256_xor_si256(d, s),
508 );
509 }
510 let rem = chunks * 32;
511 for i in rem..len {
512 dst[i] ^= src[i];
513 }
514}
515
516#[cfg(target_arch = "x86_64")]
521#[target_feature(enable = "ssse3")]
522unsafe fn mul_add_buffer_ssse3(dst: &mut [u8], src: &[u8], constant: u16) {
523 use std::arch::x86_64::*;
524
525 let tables = GfMulTables::new(constant);
526 let nibble_mask = _mm_set1_epi8(0x0F);
527
528 let tbl_lo_lo = _mm_loadu_si128(tables.lo_lo.as_ptr() as *const __m128i);
529 let tbl_lo_hi = _mm_loadu_si128(tables.lo_hi.as_ptr() as *const __m128i);
530 let tbl_hi_lo = _mm_loadu_si128(tables.hi_lo.as_ptr() as *const __m128i);
531 let tbl_hi_hi = _mm_loadu_si128(tables.hi_hi.as_ptr() as *const __m128i);
532 let tbl_ulo_lo = _mm_loadu_si128(tables.ulo_lo.as_ptr() as *const __m128i);
533 let tbl_ulo_hi = _mm_loadu_si128(tables.ulo_hi.as_ptr() as *const __m128i);
534 let tbl_uhi_lo = _mm_loadu_si128(tables.uhi_lo.as_ptr() as *const __m128i);
535 let tbl_uhi_hi = _mm_loadu_si128(tables.uhi_hi.as_ptr() as *const __m128i);
536
537 let deint_lo = _mm_setr_epi8(0, 2, 4, 6, 8, 10, 12, 14, -1, -1, -1, -1, -1, -1, -1, -1);
538 let deint_hi = _mm_setr_epi8(1, 3, 5, 7, 9, 11, 13, 15, -1, -1, -1, -1, -1, -1, -1, -1);
539
540 let len = dst.len();
541 let chunks = len / 16;
542
543 for chunk in 0..chunks {
544 let off = chunk * 16;
545 let src_data = _mm_loadu_si128(src[off..].as_ptr() as *const __m128i);
546
547 let src_lo_bytes = _mm_shuffle_epi8(src_data, deint_lo);
548 let src_hi_bytes = _mm_shuffle_epi8(src_data, deint_hi);
549
550 let lo_nib = _mm_and_si128(src_lo_bytes, nibble_mask);
551 let hi_nib = _mm_and_si128(_mm_srli_epi16(src_lo_bytes, 4), nibble_mask);
552 let ulo_nib = _mm_and_si128(src_hi_bytes, nibble_mask);
553 let uhi_nib = _mm_and_si128(_mm_srli_epi16(src_hi_bytes, 4), nibble_mask);
554
555 let r_lo = _mm_xor_si128(
556 _mm_xor_si128(
557 _mm_shuffle_epi8(tbl_lo_lo, lo_nib),
558 _mm_shuffle_epi8(tbl_hi_lo, hi_nib),
559 ),
560 _mm_xor_si128(
561 _mm_shuffle_epi8(tbl_ulo_lo, ulo_nib),
562 _mm_shuffle_epi8(tbl_uhi_lo, uhi_nib),
563 ),
564 );
565 let r_hi = _mm_xor_si128(
566 _mm_xor_si128(
567 _mm_shuffle_epi8(tbl_lo_hi, lo_nib),
568 _mm_shuffle_epi8(tbl_hi_hi, hi_nib),
569 ),
570 _mm_xor_si128(
571 _mm_shuffle_epi8(tbl_ulo_hi, ulo_nib),
572 _mm_shuffle_epi8(tbl_uhi_hi, uhi_nib),
573 ),
574 );
575
576 let result = _mm_unpacklo_epi8(r_lo, r_hi);
577 let dst_val = _mm_loadu_si128(dst[off..].as_ptr() as *const __m128i);
578 _mm_storeu_si128(
579 dst[off..].as_mut_ptr() as *mut __m128i,
580 _mm_xor_si128(dst_val, result),
581 );
582 }
583
584 let rem = chunks * 16;
585 if rem < len {
586 mul_add_buffer_scalar(&mut dst[rem..], &src[rem..], constant);
587 }
588}
589
590#[cfg(test)]
595mod tests {
596 use super::*;
597
598 #[test]
599 fn test_mul_add_buffer_scalar_basic() {
600 let src = [3u8, 0];
601 let mut dst = [0u8, 0];
602 mul_add_buffer(&mut dst, &src, 5);
603 let expected = gf::mul(3, 5);
604 let result = u16::from_le_bytes([dst[0], dst[1]]);
605 assert_eq!(result, expected);
606 }
607
608 #[test]
609 fn test_mul_add_buffer_accumulates() {
610 let src = [7u8, 0, 11, 0];
611 let mut dst = [0xFFu8, 0x00, 0x00, 0x01];
612 let constant = 42u16;
613 mul_add_buffer(&mut dst, &src, constant);
614 let expected0 = 0x00FF ^ gf::mul(constant, 7);
615 let expected1 = 0x0100 ^ gf::mul(constant, 11);
616 assert_eq!(u16::from_le_bytes([dst[0], dst[1]]), expected0);
617 assert_eq!(u16::from_le_bytes([dst[2], dst[3]]), expected1);
618 }
619
620 #[test]
621 fn test_mul_add_buffer_large() {
622 let n = 4096; let mut src = vec![0u8; n];
624 let mut dst_ref = vec![0u8; n];
625 let mut dst_simd = vec![0u8; n];
626 let constant = 12345u16;
627
628 for i in 0..n / 2 {
629 let val = (i as u16).wrapping_mul(7).wrapping_add(13);
630 src[i * 2] = val as u8;
631 src[i * 2 + 1] = (val >> 8) as u8;
632 }
633
634 mul_add_buffer_scalar(&mut dst_ref, &src, constant);
635 mul_add_buffer(&mut dst_simd, &src, constant);
636 assert_eq!(dst_simd, dst_ref, "SIMD and scalar results must match");
637 }
638
639 #[test]
640 fn test_mul_add_multi_matches_sequential() {
641 let n = 2048;
642 let src1: Vec<u8> = (0..n).map(|i| (i * 3) as u8).collect();
643 let src2: Vec<u8> = (0..n).map(|i| (i * 7 + 1) as u8).collect();
644 let src3: Vec<u8> = (0..n).map(|i| (i * 11 + 5) as u8).collect();
645 let coeffs = [100u16, 200, 300];
646 let srcs: Vec<&[u8]> = vec![&src1, &src2, &src3];
647
648 let mut dst_seq = vec![0u8; n];
650 mul_add_buffer(&mut dst_seq, &src1, 100);
651 mul_add_buffer(&mut dst_seq, &src2, 200);
652 mul_add_buffer(&mut dst_seq, &src3, 300);
653
654 let mut dst_batch = vec![0u8; n];
656 mul_add_multi(&mut dst_batch, &srcs, &coeffs);
657
658 assert_eq!(
659 dst_batch, dst_seq,
660 "Batched multi-source must match sequential"
661 );
662 }
663
664 #[test]
665 fn test_xor_buffers() {
666 let src = vec![0xAAu8; 128];
667 let mut dst = vec![0x55u8; 128];
668 xor_buffers(&mut dst, &src);
669 assert!(dst.iter().all(|&b| b == 0xFF));
670 }
671
672 #[test]
673 fn test_mul_by_zero() {
674 let src = vec![0xFF; 64];
675 let mut dst = vec![0x00; 64];
676 mul_add_buffer(&mut dst, &src, 0);
677 assert!(dst.iter().all(|&b| b == 0));
678 }
679
680 #[test]
681 fn test_mul_by_one() {
682 let src = vec![42u8, 0, 99, 0];
683 let mut dst = vec![0u8; 4];
684 mul_add_buffer(&mut dst, &src, 1);
685 assert_eq!(dst, src);
686 }
687}