1use std::{
5 fmt,
6 ops::{Add, BitAnd, BitAndAssign, BitOr, BitOrAssign, BitXor, BitXorAssign, Not, Shl, Shr},
7};
8
9use aes::cipher::{self, array::sizes};
10use bytemuck::{Pod, Zeroable};
11use rand::{Rng, distr::StandardUniform, prelude::Distribution};
12use serde::{Deserialize, Serialize};
13use subtle::{Choice, ConditionallySelectable, ConstantTimeEq};
14use thiserror::Error;
15use wide::u8x16;
16
17use crate::random_oracle::{self, RandomOracle};
18
19pub mod gf128;
20
21#[derive(Debug, Clone, Copy, Serialize, Deserialize, Default, Pod, Zeroable)]
23#[repr(transparent)]
24pub struct Block(u8x16);
25
26impl Block {
27 pub const ZERO: Self = Self(u8x16::ZERO);
29 pub const ONES: Self = Self(u8x16::MAX);
31 pub const ONE: Self = Self::new(1_u128.to_ne_bytes());
33 pub const MASK_LSB: Self = Self::pack(u64::MAX << 1, u64::MAX);
41
42 pub const BYTES: usize = 16;
44 pub const BITS: usize = 128;
46
47 #[inline]
49 pub const fn new(bytes: [u8; 16]) -> Self {
50 Self(u8x16::new(bytes))
51 }
52
53 #[inline]
55 pub const fn splat(byte: u8) -> Self {
56 Self::new([byte; 16])
57 }
58
59 #[inline]
64 pub const fn pack(low: u64, high: u64) -> Self {
65 let mut bytes = [0; 16];
66 let low = low.to_ne_bytes();
67 let mut i = 0;
68 while i < low.len() {
69 bytes[i] = low[i];
70 i += 1;
71 }
72
73 let high = high.to_ne_bytes();
74 let mut i = 0;
75 while i < high.len() {
76 bytes[i + 8] = high[i];
77 i += 1;
78 }
79
80 Self::new(bytes)
81 }
82
83 #[inline]
85 pub fn as_bytes(&self) -> &[u8; 16] {
86 self.0.as_array()
87 }
88
89 #[inline]
91 pub fn as_mut_bytes(&mut self) -> &mut [u8; 16] {
92 self.0.as_mut_array()
93 }
94
95 #[inline]
97 pub fn ro_hash(&self) -> random_oracle::Hash {
98 let mut ro = RandomOracle::new();
99 ro.update(self.as_bytes());
100 ro.finalize()
101 }
102
103 #[inline]
108 pub fn from_choices(choices: &[Choice]) -> Self {
109 assert_eq!(128, choices.len(), "choices.len() must be 128");
110 let mut bytes = [0_u8; 16];
111 let (chunks, rest) = choices.as_chunks::<8>();
112 debug_assert!(rest.is_empty());
113 for (chunk, byte) in chunks.iter().zip(&mut bytes) {
114 for (i, choice) in chunk.iter().enumerate() {
115 *byte ^= choice.unwrap_u8() << i;
116 }
117 }
118 Self::new(bytes)
119 }
120
121 #[inline]
123 pub fn low(&self) -> u64 {
124 u64::from_ne_bytes(self.as_bytes()[..8].try_into().expect("correct len"))
125 }
126
127 #[inline]
129 pub fn high(&self) -> u64 {
130 u64::from_ne_bytes(self.as_bytes()[8..].try_into().expect("correct len"))
131 }
132
133 #[inline]
135 pub fn lsb(&self) -> bool {
136 *self & Block::ONE == Block::ONE
137 }
138
139 #[inline]
141 pub fn bits(&self) -> impl Iterator<Item = bool> {
142 struct BitIter {
143 blk: Block,
144 idx: usize,
145 }
146 impl Iterator for BitIter {
147 type Item = bool;
148
149 #[inline]
150 fn next(&mut self) -> Option<Self::Item> {
151 if self.idx < Block::BITS {
152 self.idx += 1;
153 let bit = (self.blk >> (self.idx - 1)) & Block::ONE != Block::ZERO;
154 Some(bit)
155 } else {
156 None
157 }
158 }
159 }
160 BitIter { blk: *self, idx: 0 }
161 }
162}
163
164impl BitAnd for Block {
166 type Output = Self;
167
168 #[inline]
169 fn bitand(self, rhs: Self) -> Self {
170 Self(self.0 & rhs.0)
171 }
172}
173
174impl BitAndAssign for Block {
175 #[inline]
176 fn bitand_assign(&mut self, rhs: Self) {
177 *self = *self & rhs;
178 }
179}
180
181impl BitOr for Block {
182 type Output = Self;
183
184 #[inline]
185 fn bitor(self, rhs: Self) -> Self {
186 Self(self.0 | rhs.0)
187 }
188}
189
190impl BitOrAssign for Block {
191 #[inline]
192 fn bitor_assign(&mut self, rhs: Self) {
193 *self = *self | rhs;
194 }
195}
196
197impl BitXor for Block {
198 type Output = Self;
199
200 #[inline]
201 fn bitxor(self, rhs: Self) -> Self {
202 Self(self.0 ^ rhs.0)
203 }
204}
205
206impl BitXorAssign for Block {
207 #[inline]
208 fn bitxor_assign(&mut self, rhs: Self) {
209 *self = *self ^ rhs;
210 }
211}
212
213impl<Rhs> Shl<Rhs> for Block
214where
215 u128: Shl<Rhs, Output = u128>,
216{
217 type Output = Block;
218
219 #[inline]
220 fn shl(self, rhs: Rhs) -> Self::Output {
221 Self::from(u128::from(self) << rhs)
222 }
223}
224
225impl<Rhs> Shr<Rhs> for Block
226where
227 u128: Shr<Rhs, Output = u128>,
228{
229 type Output = Block;
230
231 #[inline]
232 fn shr(self, rhs: Rhs) -> Self::Output {
233 Self::from(u128::from(self) >> rhs)
234 }
235}
236
237impl Not for Block {
238 type Output = Self;
239
240 #[inline]
241 fn not(self) -> Self {
242 Self(!self.0)
243 }
244}
245
246impl PartialEq for Block {
247 fn eq(&self, other: &Self) -> bool {
248 let a: u128 = (*self).into();
249 let b: u128 = (*other).into();
250 a.ct_eq(&b).into()
251 }
252}
253
254impl Eq for Block {}
255
256impl Distribution<Block> for StandardUniform {
257 #[inline]
258 fn sample<R: Rng + ?Sized>(&self, rng: &mut R) -> Block {
259 let mut bytes = [0; 16];
260 rng.fill_bytes(&mut bytes);
261 Block::new(bytes)
262 }
263}
264
265impl AsRef<[u8]> for Block {
266 fn as_ref(&self) -> &[u8] {
267 self.as_bytes()
268 }
269}
270
271impl AsMut<[u8]> for Block {
272 #[inline]
273 fn as_mut(&mut self) -> &mut [u8] {
274 self.as_mut_bytes()
275 }
276}
277
278impl From<Block> for cipher::Array<u8, sizes::U16> {
279 #[inline]
280 fn from(value: Block) -> Self {
281 Self(*value.as_bytes())
282 }
283}
284
285impl From<cipher::Array<u8, sizes::U16>> for Block {
286 #[inline]
287 fn from(value: cipher::Array<u8, sizes::U16>) -> Self {
288 Self::new(value.0)
289 }
290}
291
292impl From<[u64; 2]> for Block {
293 #[inline]
294 fn from(value: [u64; 2]) -> Self {
295 bytemuck::cast(value)
296 }
297}
298
299impl From<Block> for [u64; 2] {
300 #[inline]
301 fn from(value: Block) -> Self {
302 bytemuck::cast(value)
303 }
304}
305
306impl From<Block> for u128 {
307 #[inline]
308 fn from(value: Block) -> Self {
309 u128::from_ne_bytes(*value.as_bytes())
311 }
312}
313
314impl From<&Block> for u128 {
315 #[inline]
316 fn from(value: &Block) -> Self {
317 u128::from_ne_bytes(*value.as_bytes())
319 }
320}
321
322impl From<usize> for Block {
323 fn from(value: usize) -> Self {
324 (value as u128).into()
325 }
326}
327
328impl From<u128> for Block {
329 #[inline]
330 fn from(value: u128) -> Self {
331 Self::new(value.to_ne_bytes())
332 }
333}
334
335impl From<&u128> for Block {
336 #[inline]
337 fn from(value: &u128) -> Self {
338 Self::new(value.to_ne_bytes())
339 }
340}
341
342#[derive(Debug, Error)]
343#[error("slice must have length of 16")]
344pub struct WrongLength;
345
346impl TryFrom<&[u8]> for Block {
347 type Error = WrongLength;
348
349 #[inline]
350 fn try_from(value: &[u8]) -> Result<Self, Self::Error> {
351 let arr = value.try_into().map_err(|_| WrongLength)?;
352 Ok(Self::new(arr))
353 }
354}
355
356#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
357mod from_arch_impls {
358 #[cfg(target_arch = "x86")]
359 use std::arch::x86::*;
360 #[cfg(target_arch = "x86_64")]
361 use std::arch::x86_64::*;
362
363 use super::Block;
364
365 impl From<__m128i> for Block {
366 #[inline]
367 fn from(value: __m128i) -> Self {
368 bytemuck::must_cast(value)
369 }
370 }
371
372 impl From<&__m128i> for Block {
373 #[inline]
374 fn from(value: &__m128i) -> Self {
375 bytemuck::must_cast(*value)
376 }
377 }
378
379 impl From<Block> for __m128i {
380 #[inline]
381 fn from(value: Block) -> Self {
382 bytemuck::must_cast(value)
383 }
384 }
385
386 impl From<&Block> for __m128i {
387 #[inline]
388 fn from(value: &Block) -> Self {
389 bytemuck::must_cast(*value)
390 }
391 }
392}
393
394impl ConditionallySelectable for Block {
395 #[inline]
396 fn conditional_select(a: &Self, b: &Self, choice: Choice) -> Self {
398 let mask = Block::new((-(choice.unwrap_u8() as i128)).to_le_bytes());
401 *a ^ (mask & (*a ^ *b))
402 }
403}
404
405impl Add for Block {
406 type Output = Block;
407
408 #[inline]
409 fn add(self, rhs: Self) -> Self::Output {
410 let a: u128 = self.into();
412 let b: u128 = rhs.into();
413 Self::from(a.wrapping_add(b))
414 }
415}
416
417impl fmt::Binary for Block {
418 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
419 fmt::Binary::fmt(&u128::from(*self), f)
420 }
421}
422
423#[cfg(feature = "num-traits")]
424impl num_traits::Zero for Block {
425 fn zero() -> Self {
426 Self::ZERO
427 }
428
429 fn is_zero(&self) -> bool {
430 *self == Self::ZERO
431 }
432}
433
434#[cfg(test)]
435mod tests {
436 use subtle::{Choice, ConditionallySelectable};
437
438 use crate::Block;
439
440 #[test]
441 fn test_block_cond_select() {
442 let choice = Choice::from(0);
443 assert_eq!(
444 Block::ZERO,
445 Block::conditional_select(&Block::ZERO, &Block::ONES, choice)
446 );
447 let choice = Choice::from(1);
448 assert_eq!(
449 Block::ONES,
450 Block::conditional_select(&Block::ZERO, &Block::ONES, choice)
451 );
452 }
453
454 #[test]
455 fn test_block_low_high() {
456 let b = Block::from(1_u128);
457 assert_eq!(1, b.low());
458 assert_eq!(0, b.high());
459 }
460
461 #[test]
462 fn test_from_into_u64_arr() {
463 let b = Block::from([42, 65]);
464 assert_eq!(42, b.low());
465 assert_eq!(65, b.high());
466 assert_eq!([42, 65], <[u64; 2]>::from(b));
467 }
468
469 #[test]
470 fn test_pack() {
471 let b = Block::pack(42, 123);
472 assert_eq!(42, b.low());
473 assert_eq!(123, b.high());
474 }
475
476 #[test]
477 fn test_mask_lsb() {
478 assert_eq!(Block::ONES ^ Block::ONE, Block::MASK_LSB);
479 }
480
481 #[test]
482 fn test_bits() {
483 let b: Block = 0b101_u128.into();
484 let mut iter = b.bits();
485 assert_eq!(Some(true), iter.next());
486 assert_eq!(Some(false), iter.next());
487 assert_eq!(Some(true), iter.next());
488 for rest in iter {
489 assert_eq!(false, rest);
490 }
491 }
492
493 #[test]
494 fn test_from_choices() {
495 let mut choices = vec![Choice::from(0); 128];
496 choices[2] = Choice::from(1);
497 choices[16] = Choice::from(1);
498 let blk = Block::from_choices(&choices);
499 assert_eq!(Block::from(1_u128 << 2 | 1_u128 << 16), blk);
500 }
501}