1use crate::{Bit, Block8, Block16, Block32, Block64, Block128};
20use crate::{
21 CanonicalDeserialize, CanonicalSerialize, Flat, FlatPromote, HardwareField, PackableField,
22 PackedFlat, TowerField,
23};
24use core::ops::{Add, AddAssign, Mul, MulAssign, Sub, SubAssign};
25use serde::{Deserialize, Serialize};
26use zeroize::Zeroize;
27
28const TAU_FLAT: u128 = 0x66340c45203fe3685d08f8c248334a81;
31
32#[derive(Copy, Clone, Default, Debug, Eq, PartialEq, Serialize, Deserialize, Zeroize)]
33#[repr(C, align(32))]
34pub struct Block256(pub [u128; 2]); impl Block256 {
37 const TAU: Self = Block256([0, 0x2000_0000_0000_0000_0000_0000_0000_0000]);
38
39 pub fn new(lo: Block128, hi: Block128) -> Self {
40 Self([lo.0, hi.0])
41 }
42
43 #[inline(always)]
44 pub fn split(self) -> (Block128, Block128) {
45 (Block128(self.0[0]), Block128(self.0[1]))
46 }
47}
48
49impl TowerField for Block256 {
50 const BITS: usize = 256;
51 const ZERO: Self = Block256([0, 0]);
52 const ONE: Self = Block256([1, 0]);
53
54 const EXTENSION_TAU: Self = Self::TAU;
55
56 fn invert(&self) -> Self {
57 let (l, h) = self.split();
58 let h2 = h * h;
59 let l2 = l * l;
60 let hl = h * l;
61 let norm = (h2 * Block128::EXTENSION_TAU) + hl + l2;
62
63 let norm_inv = norm.invert();
64 let res_hi = h * norm_inv;
65 let res_lo = (h + l) * norm_inv;
66
67 Self::new(res_lo, res_hi)
68 }
69
70 fn from_uniform_bytes(bytes: &[u8; 32]) -> Self {
71 let mut lo_buf = [0u8; 16];
72 let mut hi_buf = [0u8; 16];
73
74 lo_buf.copy_from_slice(&bytes[0..16]);
75 hi_buf.copy_from_slice(&bytes[16..32]);
76
77 Self([u128::from_le_bytes(lo_buf), u128::from_le_bytes(hi_buf)])
78 }
79}
80
81impl Add for Block256 {
82 type Output = Self;
83
84 fn add(self, rhs: Self) -> Self {
85 Self([self.0[0] ^ rhs.0[0], self.0[1] ^ rhs.0[1]])
86 }
87}
88
89impl Sub for Block256 {
90 type Output = Self;
91
92 fn sub(self, rhs: Self) -> Self {
93 self.add(rhs)
94 }
95}
96
97impl Mul for Block256 {
98 type Output = Self;
99
100 fn mul(self, rhs: Self) -> Self {
101 let (a0, a1) = self.split();
102 let (b0, b1) = rhs.split();
103
104 let v0 = a0 * b0;
105 let v1 = a1 * b1;
106 let v_sum = (a0 + a1) * (b0 + b1);
107
108 let c_hi = v0 + v_sum;
109 let c_lo = v0 + (v1 * Block128::EXTENSION_TAU);
110
111 Self::new(c_lo, c_hi)
112 }
113}
114
115impl AddAssign for Block256 {
116 fn add_assign(&mut self, rhs: Self) {
117 self.0[0] ^= rhs.0[0];
118 self.0[1] ^= rhs.0[1];
119 }
120}
121
122impl SubAssign for Block256 {
123 fn sub_assign(&mut self, rhs: Self) {
124 self.0[0] ^= rhs.0[0];
125 self.0[1] ^= rhs.0[1];
126 }
127}
128
129impl MulAssign for Block256 {
130 fn mul_assign(&mut self, rhs: Self) {
131 *self = *self * rhs;
132 }
133}
134
135impl CanonicalSerialize for Block256 {
136 fn serialized_size(&self) -> usize {
137 32
138 }
139
140 fn serialize(&self, writer: &mut [u8]) -> Result<(), ()> {
141 if writer.len() < 32 {
142 return Err(());
143 }
144
145 writer[0..16].copy_from_slice(&self.0[0].to_le_bytes());
146 writer[16..32].copy_from_slice(&self.0[1].to_le_bytes());
147
148 Ok(())
149 }
150}
151
152impl CanonicalDeserialize for Block256 {
153 fn deserialize(bytes: &[u8]) -> Result<Self, ()> {
154 if bytes.len() < 32 {
155 return Err(());
156 }
157
158 let mut lo_buf = [0u8; 16];
159 let mut hi_buf = [0u8; 16];
160
161 lo_buf.copy_from_slice(&bytes[0..16]);
162 hi_buf.copy_from_slice(&bytes[16..32]);
163
164 Ok(Self([
165 u128::from_le_bytes(lo_buf),
166 u128::from_le_bytes(hi_buf),
167 ]))
168 }
169}
170
171impl From<u8> for Block256 {
172 fn from(val: u8) -> Self {
173 Self([val as u128, 0])
174 }
175}
176
177impl From<u32> for Block256 {
178 #[inline]
179 fn from(val: u32) -> Self {
180 Self([val as u128, 0])
181 }
182}
183
184impl From<u64> for Block256 {
185 #[inline]
186 fn from(val: u64) -> Self {
187 Self([val as u128, 0])
188 }
189}
190
191impl From<u128> for Block256 {
192 #[inline]
193 fn from(val: u128) -> Self {
194 Self([val, 0])
195 }
196}
197
198impl From<Bit> for Block256 {
199 #[inline(always)]
200 fn from(val: Bit) -> Self {
201 Self([val.get() as u128, 0])
202 }
203}
204
205impl From<Block8> for Block256 {
206 #[inline(always)]
207 fn from(val: Block8) -> Self {
208 Self([val.0 as u128, 0])
209 }
210}
211
212impl From<Block16> for Block256 {
213 #[inline(always)]
214 fn from(val: Block16) -> Self {
215 Self([val.0 as u128, 0])
216 }
217}
218
219impl From<Block32> for Block256 {
220 #[inline(always)]
221 fn from(val: Block32) -> Self {
222 Self([val.0 as u128, 0])
223 }
224}
225
226impl From<Block64> for Block256 {
227 #[inline(always)]
228 fn from(val: Block64) -> Self {
229 Self([val.0 as u128, 0])
230 }
231}
232
233impl From<Block128> for Block256 {
234 #[inline(always)]
235 fn from(val: Block128) -> Self {
236 Self([val.0, 0])
237 }
238}
239
240pub const PACKED_WIDTH_256: usize = 2;
245
246#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
247#[repr(C, align(64))]
248pub struct PackedBlock256(pub [Block256; PACKED_WIDTH_256]);
249
250impl PackedBlock256 {
251 #[inline(always)]
252 pub fn zero() -> Self {
253 Self([Block256::ZERO; PACKED_WIDTH_256])
254 }
255
256 #[inline(always)]
257 pub fn broadcast(val: Block256) -> Self {
258 Self([val; PACKED_WIDTH_256])
259 }
260}
261
262impl PackableField for Block256 {
263 type Packed = PackedBlock256;
264
265 const WIDTH: usize = PACKED_WIDTH_256;
266
267 #[inline(always)]
268 fn pack(chunk: &[Self]) -> Self::Packed {
269 assert!(
270 chunk.len() >= PACKED_WIDTH_256,
271 "PackableField::pack: input slice too short",
272 );
273
274 let mut arr = [Self::ZERO; PACKED_WIDTH_256];
275 arr.copy_from_slice(&chunk[..PACKED_WIDTH_256]);
276
277 PackedBlock256(arr)
278 }
279
280 #[inline(always)]
281 fn unpack(packed: Self::Packed, output: &mut [Self]) {
282 assert!(
283 output.len() >= PACKED_WIDTH_256,
284 "PackableField::unpack: output slice too short",
285 );
286
287 output[..PACKED_WIDTH_256].copy_from_slice(&packed.0);
288 }
289}
290
291impl Add for PackedBlock256 {
292 type Output = Self;
293
294 #[inline(always)]
295 fn add(self, rhs: Self) -> Self {
296 let mut res = [Block256::ZERO; PACKED_WIDTH_256];
297 for ((out, l), r) in res.iter_mut().zip(self.0.iter()).zip(rhs.0.iter()) {
298 *out = *l + *r;
299 }
300
301 Self(res)
302 }
303}
304
305impl AddAssign for PackedBlock256 {
306 #[inline(always)]
307 fn add_assign(&mut self, rhs: Self) {
308 for (l, r) in self.0.iter_mut().zip(rhs.0.iter()) {
309 *l += *r;
310 }
311 }
312}
313
314impl Sub for PackedBlock256 {
315 type Output = Self;
316
317 #[inline(always)]
318 fn sub(self, rhs: Self) -> Self {
319 self.add(rhs)
320 }
321}
322
323impl SubAssign for PackedBlock256 {
324 #[inline(always)]
325 fn sub_assign(&mut self, rhs: Self) {
326 self.add_assign(rhs);
327 }
328}
329
330impl Mul for PackedBlock256 {
331 type Output = Self;
332
333 #[inline(always)]
334 fn mul(self, rhs: Self) -> Self {
335 let mut res = [Block256::ZERO; PACKED_WIDTH_256];
336 for ((out, l), r) in res.iter_mut().zip(self.0.iter()).zip(rhs.0.iter()) {
337 *out = *l * *r;
338 }
339
340 Self(res)
341 }
342}
343
344impl MulAssign for PackedBlock256 {
345 #[inline(always)]
346 fn mul_assign(&mut self, rhs: Self) {
347 for (l, r) in self.0.iter_mut().zip(rhs.0.iter()) {
348 *l *= *r;
349 }
350 }
351}
352
353impl Mul<Block256> for PackedBlock256 {
354 type Output = Self;
355
356 #[inline(always)]
357 fn mul(self, rhs: Block256) -> Self {
358 let mut res = [Block256::ZERO; PACKED_WIDTH_256];
359 for (out, v) in res.iter_mut().zip(self.0.iter()) {
360 *out = *v * rhs;
361 }
362
363 Self(res)
364 }
365}
366
367impl MulAssign<Block256> for PackedBlock256 {
368 #[inline(always)]
369 fn mul_assign(&mut self, rhs: Block256) {
370 for v in self.0.iter_mut() {
371 *v *= rhs;
372 }
373 }
374}
375
376impl HardwareField for Block256 {
377 #[inline(always)]
378 fn to_hardware(self) -> Flat<Self> {
379 let (lo, hi) = self.split();
380 let flat_lo = lo.to_hardware().into_raw().0;
381 let flat_hi = hi.to_hardware().into_raw().0;
382
383 Flat::from_raw(Block256([flat_lo, flat_hi]))
384 }
385
386 #[inline(always)]
387 fn from_hardware(value: Flat<Self>) -> Self {
388 let raw = value.into_raw();
389 let lo = Block128::from_hardware(Flat::from_raw(Block128(raw.0[0])));
390 let hi = Block128::from_hardware(Flat::from_raw(Block128(raw.0[1])));
391
392 Self::new(lo, hi)
393 }
394
395 #[inline(always)]
396 fn add_hardware(lhs: Flat<Self>, rhs: Flat<Self>) -> Flat<Self> {
397 let l = lhs.into_raw();
398 let r = rhs.into_raw();
399
400 let c_lo = Block128::add_hardware(
401 Flat::from_raw(Block128(l.0[0])),
402 Flat::from_raw(Block128(r.0[0])),
403 );
404 let c_hi = Block128::add_hardware(
405 Flat::from_raw(Block128(l.0[1])),
406 Flat::from_raw(Block128(r.0[1])),
407 );
408
409 Flat::from_raw(Block256([c_lo.into_raw().0, c_hi.into_raw().0]))
410 }
411
412 #[inline(always)]
413 fn add_hardware_packed(lhs: PackedFlat<Self>, rhs: PackedFlat<Self>) -> PackedFlat<Self> {
414 let lhs = lhs.into_raw().0;
415 let rhs = rhs.into_raw().0;
416
417 let mut res = [Block256::ZERO; PACKED_WIDTH_256];
418 for i in 0..PACKED_WIDTH_256 {
419 res[i] = Self::add_hardware(Flat::from_raw(lhs[i]), Flat::from_raw(rhs[i])).into_raw();
420 }
421
422 PackedFlat::from_raw(PackedBlock256(res))
423 }
424
425 #[inline(always)]
426 fn mul_hardware(lhs: Flat<Self>, rhs: Flat<Self>) -> Flat<Self> {
427 let a_lo = Flat::from_raw(Block128(lhs.into_raw().0[0]));
428 let a_hi = Flat::from_raw(Block128(lhs.into_raw().0[1]));
429 let b_lo = Flat::from_raw(Block128(rhs.into_raw().0[0]));
430 let b_hi = Flat::from_raw(Block128(rhs.into_raw().0[1]));
431
432 let tau = Flat::from_raw(Block128(TAU_FLAT));
433
434 let v0 = Block128::mul_hardware(a_lo, b_lo);
435 let v1 = Block128::mul_hardware(a_hi, b_hi);
436
437 let a_sum = Block128::add_hardware(a_lo, a_hi);
438 let b_sum = Block128::add_hardware(b_lo, b_hi);
439 let v_sum = Block128::mul_hardware(a_sum, b_sum);
440
441 let c_hi = Block128::add_hardware(v0, v_sum);
442
443 let v1_tau = Block128::mul_hardware(v1, tau);
444 let c_lo = Block128::add_hardware(v0, v1_tau);
445
446 Flat::from_raw(Block256([c_lo.into_raw().0, c_hi.into_raw().0]))
447 }
448
449 #[inline(always)]
450 fn mul_hardware_packed(lhs: PackedFlat<Self>, rhs: PackedFlat<Self>) -> PackedFlat<Self> {
451 let lhs = lhs.into_raw().0;
452 let rhs = rhs.into_raw().0;
453
454 let mut res = [Block256::ZERO; PACKED_WIDTH_256];
455 for i in 0..PACKED_WIDTH_256 {
456 res[i] = Self::mul_hardware(Flat::from_raw(lhs[i]), Flat::from_raw(rhs[i])).into_raw();
457 }
458
459 PackedFlat::from_raw(PackedBlock256(res))
460 }
461
462 #[inline(always)]
463 fn mul_hardware_scalar_packed(lhs: PackedFlat<Self>, rhs: Flat<Self>) -> PackedFlat<Self> {
464 let broadcasted = PackedBlock256::broadcast(rhs.into_raw());
465 Self::mul_hardware_packed(lhs, PackedFlat::from_raw(broadcasted))
466 }
467
468 #[inline(always)]
469 fn tower_bit_from_hardware(value: Flat<Self>, bit_idx: usize) -> u8 {
470 if bit_idx < 128 {
471 Block128::tower_bit_from_hardware(
472 Flat::from_raw(Block128(value.into_raw().0[0])),
473 bit_idx,
474 )
475 } else {
476 Block128::tower_bit_from_hardware(
477 Flat::from_raw(Block128(value.into_raw().0[1])),
478 bit_idx - 128,
479 )
480 }
481 }
482}
483
484const PROMOTE_CHUNK: usize = 64;
485
486impl FlatPromote<Block8> for Block256 {
487 #[inline(always)]
488 fn promote_flat(val: Flat<Block8>) -> Flat<Self> {
489 let promoted = Block128::promote_flat(val);
490 Flat::from_raw(Block256([promoted.into_raw().0, 0]))
491 }
492
493 fn promote_flat_batch(input: &[Flat<Block8>], output: &mut [Flat<Self>]) {
494 promote_chunked(input, output);
495 }
496}
497
498impl FlatPromote<Block16> for Block256 {
499 #[inline(always)]
500 fn promote_flat(val: Flat<Block16>) -> Flat<Self> {
501 let promoted = Block128::promote_flat(val);
502 Flat::from_raw(Block256([promoted.into_raw().0, 0]))
503 }
504
505 fn promote_flat_batch(input: &[Flat<Block16>], output: &mut [Flat<Self>]) {
506 promote_chunked(input, output);
507 }
508}
509
510impl FlatPromote<Block32> for Block256 {
511 #[inline(always)]
512 fn promote_flat(val: Flat<Block32>) -> Flat<Self> {
513 let promoted = Block128::promote_flat(val);
514 Flat::from_raw(Block256([promoted.into_raw().0, 0]))
515 }
516
517 fn promote_flat_batch(input: &[Flat<Block32>], output: &mut [Flat<Self>]) {
518 promote_chunked(input, output);
519 }
520}
521
522impl FlatPromote<Block64> for Block256 {
523 #[inline(always)]
524 fn promote_flat(val: Flat<Block64>) -> Flat<Self> {
525 let promoted = Block128::promote_flat(val);
526 Flat::from_raw(Block256([promoted.into_raw().0, 0]))
527 }
528
529 fn promote_flat_batch(input: &[Flat<Block64>], output: &mut [Flat<Self>]) {
530 promote_chunked(input, output);
531 }
532}
533
534impl FlatPromote<Block128> for Block256 {
535 #[inline(always)]
536 fn promote_flat(val: Flat<Block128>) -> Flat<Self> {
537 Flat::from_raw(Block256([val.into_raw().0, 0]))
538 }
539
540 fn promote_flat_batch(input: &[Flat<Block128>], output: &mut [Flat<Self>]) {
541 let n = input.len().min(output.len());
542 for i in 0..n {
543 output[i] = Flat::from_raw(Block256([input[i].into_raw().0, 0]));
544 }
545 }
546}
547
548#[inline(always)]
549fn promote_chunked<FromF>(input: &[Flat<FromF>], output: &mut [Flat<Block256>])
550where
551 FromF: HardwareField,
552 Block128: FlatPromote<FromF>,
553{
554 let n = input.len().min(output.len());
555
556 let mut scratch = [Flat::from_raw(Block128::ZERO); PROMOTE_CHUNK];
557 let mut i = 0;
558
559 while i < n {
560 let len = (n - i).min(PROMOTE_CHUNK);
561 Block128::promote_flat_batch(&input[i..i + len], &mut scratch[..len]);
562
563 for j in 0..len {
564 output[i + j] = Flat::from_raw(Block256([scratch[j].into_raw().0, 0]));
565 }
566
567 i += len;
568 }
569}
570
571#[cfg(test)]
572mod tests {
573 use super::*;
574 use rand::{RngExt, rng};
575
576 #[test]
577 fn tau_flat_matches_derived() {
578 let derived = Block128::EXTENSION_TAU.to_hardware().into_raw().0;
579 assert_eq!(
580 TAU_FLAT, derived,
581 "TAU_FLAT drifted from Block128::EXTENSION_TAU.to_hardware()",
582 );
583 }
584
585 #[test]
590 fn tower_constants() {
591 let tau256 = Block256::EXTENSION_TAU;
594 let (lo256, hi256) = tau256.split();
595 assert_eq!(lo256, Block128::ZERO);
596 assert_eq!(hi256, Block128::EXTENSION_TAU);
597 }
598
599 #[test]
600 fn add_truth() {
601 let zero = Block256::ZERO;
602 let one = Block256::ONE;
603
604 assert_eq!(zero + zero, zero);
605 assert_eq!(zero + one, one);
606 assert_eq!(one + zero, one);
607 assert_eq!(one + one, zero);
608 }
609
610 #[test]
611 fn mul_truth() {
612 let zero = Block256::ZERO;
613 let one = Block256::ONE;
614
615 assert_eq!(zero * zero, zero);
616 assert_eq!(zero * one, zero);
617 assert_eq!(one * one, one);
618 }
619
620 #[test]
621 fn add() {
622 assert_eq!(Block256([5, 0]) + Block256([3, 0]), Block256([6, 0]));
625 }
626
627 #[test]
628 fn mul_simple() {
629 assert_eq!(
631 Block256::from(2u32) * Block256::from(2u32),
632 Block256::from(4u32)
633 );
634 }
635
636 #[test]
637 fn mul_overflow() {
638 assert_eq!(
640 Block256::from(0x57u32) * Block256::from(0x83u32),
641 Block256::from(0xC1u32)
642 );
643 }
644
645 #[test]
646 fn karatsuba_correctness() {
647 let y = Block256::new(Block128::ZERO, Block128::ONE);
652 let squared = y * y;
653
654 let (res_lo, res_hi) = squared.split();
655
656 assert_eq!(res_hi, Block128::ONE, "Y^2 should contain Y component");
657 assert_eq!(
658 res_lo,
659 Block128::EXTENSION_TAU,
660 "Y^2 should contain tau_256 component"
661 );
662 }
663
664 #[test]
665 fn security_zeroize() {
666 let mut secret_val = Block256([0xDEAD_BEEF_CAFE_BABE_u128, 0xFEED_FACE_BAAD_F00D_u128]);
667 assert_ne!(secret_val, Block256::ZERO);
668
669 secret_val.zeroize();
670
671 assert_eq!(secret_val, Block256::ZERO, "Memory was not wiped!");
672 assert_eq!(
673 secret_val.0,
674 [0u128, 0u128],
675 "Underlying memory leak detected"
676 );
677 }
678
679 #[test]
680 fn invert_zero() {
681 assert_eq!(
682 Block256::ZERO.invert(),
683 Block256::ZERO,
684 "invert(0) must return 0"
685 );
686 }
687
688 #[test]
689 fn inversion_random() {
690 let mut rng = rng();
691 for _ in 0..1000 {
692 let val = Block256([rng.random(), rng.random()]);
693
694 if val != Block256::ZERO {
695 let inv = val.invert();
696 let identity = val * inv;
697
698 assert_eq!(
699 identity,
700 Block256::ONE,
701 "Inversion identity failed: a * a^-1 != 1"
702 );
703 }
704 }
705 }
706
707 #[test]
708 fn tower_embedding() {
709 let mut rng = rng();
710 for _ in 0..100 {
711 let a = Block128(rng.random());
712 let b = Block128(rng.random());
713
714 let a_lifted: Block256 = a.into();
717 let (lo, hi) = a_lifted.split();
718
719 assert_eq!(lo, a, "Embedding structure failed: low part mismatch");
720 assert_eq!(
721 hi,
722 Block128::ZERO,
723 "Embedding structure failed: high part must be zero"
724 );
725
726 let sum_sub = a + b;
728 let sum_lifted: Block256 = sum_sub.into();
729 let sum_in_super = Block256::from(a) + Block256::from(b);
730
731 assert_eq!(sum_lifted, sum_in_super, "Homomorphism failed: add");
732
733 let prod_sub = a * b;
735 let prod_lifted: Block256 = prod_sub.into();
736 let prod_in_super = Block256::from(a) * Block256::from(b);
737
738 assert_eq!(prod_lifted, prod_in_super, "Homomorphism failed: mul");
739 }
740 }
741
742 #[test]
747 fn isomorphism_roundtrip() {
748 let mut rng = rng();
749 for _ in 0..1000 {
750 let val = Block256([rng.random::<u128>(), rng.random::<u128>()]);
751 assert_eq!(val.to_hardware().to_tower(), val);
752 }
753 }
754
755 #[test]
756 fn flat_mul_homomorphism() {
757 let mut rng = rng();
758 for _ in 0..1000 {
759 let a = Block256([rng.random(), rng.random()]);
760 let b = Block256([rng.random(), rng.random()]);
761
762 let expected_flat = (a * b).to_hardware();
763 let actual_flat = a.to_hardware() * b.to_hardware();
764
765 assert_eq!(
766 actual_flat, expected_flat,
767 "Block256 flat multiplication mismatch: (a*b)^H != a^H * b^H"
768 );
769 }
770 }
771
772 #[test]
773 fn packed_consistency() {
774 let mut rng = rng();
775 for _ in 0..100 {
776 let mut a_vals = [Block256::ZERO; PACKED_WIDTH_256];
777 let mut b_vals = [Block256::ZERO; PACKED_WIDTH_256];
778
779 for i in 0..PACKED_WIDTH_256 {
780 a_vals[i] = Block256([rng.random::<u128>(), rng.random::<u128>()]);
781 b_vals[i] = Block256([rng.random::<u128>(), rng.random::<u128>()]);
782 }
783
784 let a_flat_vals = a_vals.map(|x| x.to_hardware());
785 let b_flat_vals = b_vals.map(|x| x.to_hardware());
786 let a_packed = Flat::<Block256>::pack(&a_flat_vals);
787 let b_packed = Flat::<Block256>::pack(&b_flat_vals);
788
789 let add_res = Block256::add_hardware_packed(a_packed, b_packed);
790
791 let mut add_out = [Block256::ZERO.to_hardware(); PACKED_WIDTH_256];
792 Flat::<Block256>::unpack(add_res, &mut add_out);
793
794 for i in 0..PACKED_WIDTH_256 {
795 assert_eq!(
796 add_out[i],
797 (a_vals[i] + b_vals[i]).to_hardware(),
798 "Block256 SIMD add mismatch at index {}",
799 i
800 );
801 }
802
803 let mul_res = Block256::mul_hardware_packed(a_packed, b_packed);
804
805 let mut mul_out = [Block256::ZERO.to_hardware(); PACKED_WIDTH_256];
806 Flat::<Block256>::unpack(mul_res, &mut mul_out);
807
808 for i in 0..PACKED_WIDTH_256 {
809 let expected_flat = (a_vals[i] * b_vals[i]).to_hardware();
810 assert_eq!(
811 mul_out[i], expected_flat,
812 "Block256 SIMD mul mismatch at index {}",
813 i
814 );
815 }
816 }
817 }
818
819 #[test]
820 fn tower_bit_from_hardware_matches_tower() {
821 let mut rng = rng();
822 for _ in 0..64 {
823 let val = Block256([rng.random::<u128>(), rng.random::<u128>()]);
824 let flat = val.to_hardware();
825
826 for bit in 0..Block256::BITS {
827 let expected = if bit < 128 {
828 ((val.0[0] >> bit) & 1) as u8
829 } else {
830 ((val.0[1] >> (bit - 128)) & 1) as u8
831 };
832
833 assert_eq!(
834 Block256::tower_bit_from_hardware(flat, bit),
835 expected,
836 "tower_bit mismatch at bit {}",
837 bit
838 );
839 }
840 }
841 }
842
843 #[test]
848 fn promote_flat_batch_matches_scalar_block8() {
849 let mut rng = rng();
850 let input: Vec<Flat<Block8>> = (0..200)
851 .map(|_| Block8(rng.random::<u8>()).to_hardware())
852 .collect();
853
854 let mut batch_out = vec![Flat::from_raw(Block256::ZERO); input.len()];
855 <Block256 as FlatPromote<Block8>>::promote_flat_batch(&input, &mut batch_out);
856
857 for i in 0..input.len() {
858 let scalar = <Block256 as FlatPromote<Block8>>::promote_flat(input[i]);
859 assert_eq!(
860 batch_out[i], scalar,
861 "Block8 batch/scalar mismatch at {}",
862 i
863 );
864 }
865 }
866
867 #[test]
868 fn promote_flat_batch_matches_scalar_block16() {
869 let mut rng = rng();
870 let input: Vec<Flat<Block16>> = (0..200)
871 .map(|_| Block16(rng.random::<u16>()).to_hardware())
872 .collect();
873
874 let mut batch_out = vec![Flat::from_raw(Block256::ZERO); input.len()];
875 <Block256 as FlatPromote<Block16>>::promote_flat_batch(&input, &mut batch_out);
876
877 for i in 0..input.len() {
878 let scalar = <Block256 as FlatPromote<Block16>>::promote_flat(input[i]);
879 assert_eq!(
880 batch_out[i], scalar,
881 "Block16 batch/scalar mismatch at {}",
882 i
883 );
884 }
885 }
886
887 #[test]
888 fn promote_flat_batch_matches_scalar_block32() {
889 let mut rng = rng();
890 let input: Vec<Flat<Block32>> = (0..200)
891 .map(|_| Block32(rng.random::<u32>()).to_hardware())
892 .collect();
893
894 let mut batch_out = vec![Flat::from_raw(Block256::ZERO); input.len()];
895 <Block256 as FlatPromote<Block32>>::promote_flat_batch(&input, &mut batch_out);
896
897 for i in 0..input.len() {
898 let scalar = <Block256 as FlatPromote<Block32>>::promote_flat(input[i]);
899 assert_eq!(
900 batch_out[i], scalar,
901 "Block32 batch/scalar mismatch at {}",
902 i
903 );
904 }
905 }
906
907 #[test]
908 fn promote_flat_batch_matches_scalar_block64() {
909 let mut rng = rng();
910 let input: Vec<Flat<Block64>> = (0..200)
911 .map(|_| Block64(rng.random::<u64>()).to_hardware())
912 .collect();
913
914 let mut batch_out = vec![Flat::from_raw(Block256::ZERO); input.len()];
915 <Block256 as FlatPromote<Block64>>::promote_flat_batch(&input, &mut batch_out);
916
917 for i in 0..input.len() {
918 let scalar = <Block256 as FlatPromote<Block64>>::promote_flat(input[i]);
919 assert_eq!(
920 batch_out[i], scalar,
921 "Block64 batch/scalar mismatch at {}",
922 i
923 );
924 }
925 }
926
927 #[test]
928 fn promote_flat_batch_matches_scalar_block128() {
929 let mut rng = rng();
930 let input: Vec<Flat<Block128>> = (0..200)
931 .map(|_| Block128(rng.random::<u128>()).to_hardware())
932 .collect();
933
934 let mut batch_out = vec![Flat::from_raw(Block256::ZERO); input.len()];
935 <Block256 as FlatPromote<Block128>>::promote_flat_batch(&input, &mut batch_out);
936
937 for i in 0..input.len() {
938 let scalar = <Block256 as FlatPromote<Block128>>::promote_flat(input[i]);
939 assert_eq!(
940 batch_out[i], scalar,
941 "Block128 batch/scalar mismatch at {}",
942 i
943 );
944 }
945 }
946
947 #[test]
948 fn promote_flat_batch_partial_slice() {
949 let mut rng = rng();
950 let input: Vec<Flat<Block8>> = (0..10)
951 .map(|_| Block8(rng.random::<u8>()).to_hardware())
952 .collect();
953
954 let mut out_short = vec![Flat::from_raw(Block256::ZERO); 5];
955 <Block256 as FlatPromote<Block8>>::promote_flat_batch(&input, &mut out_short);
956
957 for i in 0..5 {
958 let scalar = <Block256 as FlatPromote<Block8>>::promote_flat(input[i]);
959 assert_eq!(out_short[i], scalar);
960 }
961
962 let short_input = &input[..3];
963
964 let mut out_long = vec![Flat::from_raw(Block256::ZERO); 10];
965 <Block256 as FlatPromote<Block8>>::promote_flat_batch(short_input, &mut out_long);
966
967 for i in 0..3 {
968 let scalar = <Block256 as FlatPromote<Block8>>::promote_flat(short_input[i]);
969 assert_eq!(out_long[i], scalar);
970 }
971
972 for val in out_long.iter().skip(3) {
973 assert_eq!(*val, Flat::from_raw(Block256::ZERO));
974 }
975 }
976
977 #[test]
978 fn promote_flat_batch_across_chunk_boundary() {
979 let mut rng = rng();
980 for &n in &[
982 PROMOTE_CHUNK - 1,
983 PROMOTE_CHUNK,
984 PROMOTE_CHUNK + 1,
985 PROMOTE_CHUNK * 2 + 3,
986 ] {
987 let input: Vec<Flat<Block8>> = (0..n)
988 .map(|_| Block8(rng.random::<u8>()).to_hardware())
989 .collect();
990
991 let mut batch_out = vec![Flat::from_raw(Block256::ZERO); n];
992 <Block256 as FlatPromote<Block8>>::promote_flat_batch(&input, &mut batch_out);
993
994 for i in 0..n {
995 let scalar = <Block256 as FlatPromote<Block8>>::promote_flat(input[i]);
996 assert_eq!(batch_out[i], scalar, "n={}, idx={}", n, i);
997 }
998 }
999 }
1000}