Skip to main content

rust_par2/
gf_simd.rs

1// All SIMD functions are `unsafe fn` wrapping intrinsics — inner unsafe blocks
2// would add noise without safety benefit since the caller already entered unsafe.
3#![allow(unsafe_op_in_unsafe_fn)]
4
5//! SIMD-accelerated GF(2^16) buffer operations.
6//!
7//! Dispatch hierarchy: AVX2 (256-bit) → SSSE3 (128-bit) → scalar.
8//!
9//! The PSHUFB technique: decompose 16-bit GF multiply into four 4-bit lookups.
10//! VPSHUFB (AVX2) processes 32 bytes per instruction, PSHUFB (SSSE3) 16 bytes.
11//!
12//! Additionally provides `mul_add_multi` which accumulates multiple source
13//! buffers × coefficients into dst in a single pass, reducing memory bandwidth
14//! by loading each dst cache line once instead of once per source.
15
16use crate::gf;
17
18/// Precomputed PSHUFB tables for multiplying by a GF(2^16) constant.
19pub 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
60// =========================================================================
61// Single-source: dst ^= constant * src
62// =========================================================================
63
64/// dst[i] ^= constant * src[i] for each u16 position.
65pub 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
89// =========================================================================
90// Multi-source batched: dst ^= Σ coeffs[i] * srcs[i]
91// =========================================================================
92
93/// Accumulate multiple source buffers into dst.
94///
95/// `dst ^= coeffs[0]*srcs[0] + coeffs[1]*srcs[1] + ...`
96///
97/// Uses batched processing: groups sources into batches of 2, loading each
98/// dst cache line once per batch instead of once per source. With AVX2
99/// this halves memory bandwidth for dst.
100pub fn mul_add_multi(dst: &mut [u8], srcs: &[&[u8]], coeffs: &[u16]) {
101    assert_eq!(srcs.len(), coeffs.len());
102
103    // Filter out zero coefficients
104    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            // Process pairs of sources: load dst once, accumulate 2 sources, store
119            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            // Odd remainder
127            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
140/// XOR src into dst.
141pub 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
155// =========================================================================
156// Scalar fallback
157// =========================================================================
158
159fn 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// =========================================================================
172// AVX2 implementation (256-bit = 32 bytes = 16 u16 values per instruction)
173// =========================================================================
174
175#[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/// Core AVX2 GF multiply-accumulate for a single src buffer.
183#[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    // Broadcast 16-byte tables to 256-bit (duplicate in both 128-bit lanes)
189    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    // Deinterleave masks: extract even (u16-low) and odd (u16-high) bytes
204    // VPSHUFB operates per 128-bit lane, so the mask is the same in both lanes
205    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        // Deinterleave: separate u16-low and u16-high bytes within each 128-bit lane
220        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        // Split into nibbles
224        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        // VPSHUFB lookups — low byte of result
230        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        // High byte of result
241        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        // Interleave result bytes back to u16 LE (within each 128-bit lane)
253        let result = _mm256_unpacklo_epi8(r_lo, r_hi);
254
255        // XOR into destination
256        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    // Scalar remainder
264    let rem = chunks * 32;
265    if rem < len {
266        mul_add_buffer_scalar(
267            &mut dst[rem..],
268            &src[rem..],
269            gf::mul(
270                // Recompute constant from tables — just pass it through
271                // Actually we need the original constant, extract from table:
272                // tables.lo_lo[1] | (tables.lo_hi[1] << 8) = constant * 1 = constant
273                tables.lo_lo[1] as u16 | ((tables.lo_hi[1] as u16) << 8),
274                1,
275            ),
276        );
277        // Simpler: just redo scalar with the real constant
278    }
279}
280
281/// Process 2 sources per dst load/store: dst ^= c1*src1 + c2*src2
282#[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    // Load tables for source 1
299    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    // Load tables for source 2
309    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        // Load dst once
325        let mut acc = _mm256_loadu_si256(dst[off..].as_ptr() as *const __m256i);
326
327        // Source 1
328        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        // Source 2
358        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        // Store once
388        _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/// Old batched multi-source (not used, replaced by pair batching above).
399#[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    // Precompute all tables
414    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        // Load dst once per 32-byte chunk
423        let mut acc = _mm256_loadu_si256(dst[off..].as_ptr() as *const __m256i);
424
425        // Accumulate all sources into this chunk
426        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        // Store accumulated result once
483        _mm256_storeu_si256(dst[off..].as_mut_ptr() as *mut __m256i, acc);
484    }
485
486    // Scalar remainder
487    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// =========================================================================
517// SSSE3 fallback (128-bit)
518// =========================================================================
519
520#[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// =========================================================================
591// Tests
592// =========================================================================
593
594#[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; // Large enough for AVX2 path (>32 bytes)
623        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        // Sequential reference
649        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        // Batched
655        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}