Skip to main content

cryprot_core/
block.rs

1//! A 128-bit [`Block`] type.
2//!
3//! Operations on [`Block`]s will use SIMD instructions where possible.
4use 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/// A 128-bit block. Uses SIMD operations where available.
22#[derive(Debug, Clone, Copy, Serialize, Deserialize, Default, Pod, Zeroable)]
23#[repr(transparent)]
24pub struct Block(u8x16);
25
26impl Block {
27    /// All bits set to 0.
28    pub const ZERO: Self = Self(u8x16::ZERO);
29    /// All bits set to 1.
30    pub const ONES: Self = Self(u8x16::MAX);
31    /// Lsb set to 1, all others zero.
32    pub const ONE: Self = Self::new(1_u128.to_ne_bytes());
33    /// Mask to mask off the LSB of a Block.
34    /// ```rust
35    /// # use cryprot_core::Block;
36    /// let b = Block::ONES;
37    /// let masked = b & Block::MASK_LSB;
38    /// assert_eq!(masked, Block::ONES << 1)
39    /// ```
40    pub const MASK_LSB: Self = Self::pack(u64::MAX << 1, u64::MAX);
41
42    /// 16 bytes in a Block.
43    pub const BYTES: usize = 16;
44    /// 128 bits in a block.
45    pub const BITS: usize = 128;
46
47    /// Create a new block from bytes.
48    #[inline]
49    pub const fn new(bytes: [u8; 16]) -> Self {
50        Self(u8x16::new(bytes))
51    }
52
53    /// Create a block with all bytes set to `byte`.
54    #[inline]
55    pub const fn splat(byte: u8) -> Self {
56        Self::new([byte; 16])
57    }
58
59    /// Pack two `u64` into a Block. Usable in const context.
60    ///
61    /// In non-const contexts, using `Block::from([low, high])` is likely
62    /// faster.
63    #[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    /// Bytes of the block.
84    #[inline]
85    pub fn as_bytes(&self) -> &[u8; 16] {
86        self.0.as_array()
87    }
88
89    /// Mutable bytes of the block.
90    #[inline]
91    pub fn as_mut_bytes(&mut self) -> &mut [u8; 16] {
92        self.0.as_mut_array()
93    }
94
95    /// Hash the block with a [`random_oracle`].
96    #[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    ///  Create a block from 128 [`Choice`]s.
104    ///
105    /// # Panics
106    /// If choices.len() != 128
107    #[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    /// Low 64 bits of the block.
122    #[inline]
123    pub fn low(&self) -> u64 {
124        u64::from_ne_bytes(self.as_bytes()[..8].try_into().expect("correct len"))
125    }
126
127    /// High 64 bits of the block.
128    #[inline]
129    pub fn high(&self) -> u64 {
130        u64::from_ne_bytes(self.as_bytes()[8..].try_into().expect("correct len"))
131    }
132
133    /// Least significant bit of the block
134    #[inline]
135    pub fn lsb(&self) -> bool {
136        *self & Block::ONE == Block::ONE
137    }
138
139    /// Iterator over bits of the Block.
140    #[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
164// Implement standard operators for more ergonomic usage
165impl 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        // todo correct endianness?
310        u128::from_ne_bytes(*value.as_bytes())
311    }
312}
313
314impl From<&Block> for u128 {
315    #[inline]
316    fn from(value: &Block) -> Self {
317        // todo correct endianness?
318        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    // adapted from https://github.com/dalek-cryptography/subtle/blob/369e7463e85921377a5f2df80aabcbbc6d57a930/src/lib.rs#L510-L517
397    fn conditional_select(a: &Self, b: &Self, choice: Choice) -> Self {
398        // if choice = 0, mask = (-0) = 0000...0000
399        // if choice = 1, mask = (-1) = 1111...1111
400        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        // todo is this a sensible implementation?
411        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}