Skip to main content

aes/backends/x86_vaes256/
encdec.rs

1use super::RoundKeys;
2use crate::Block;
3use cipher::{Array, array::ArraySize, consts::U2, inout::InOut, typenum::Quot};
4use core::ops::Div;
5
6#[cfg(target_arch = "x86")]
7use core::arch::x86::*;
8#[cfg(target_arch = "x86_64")]
9use core::arch::x86_64::*;
10
11pub(super) type RoundKeys2<const ROUNDS: usize> = [__m256i; ROUNDS];
12
13type BatchBlocks<ParBlocks> = Array<__m256i, Quot<ParBlocks, U2>>;
14
15#[inline]
16#[target_feature(enable = "avx2")]
17pub(super) fn broadcast_keys<const RK: usize>(keys: &RoundKeys<RK>) -> RoundKeys2<RK> {
18    keys.map(|key| _mm256_broadcastsi128_si256(key))
19}
20
21#[inline]
22#[target_feature(enable = "vaes")]
23pub(super) fn batch_encrypt<const RK: usize, ParBlocks>(
24    keys: &RoundKeys2<RK>,
25    mut blocks: InOut<'_, '_, Array<Block, ParBlocks>>,
26) where
27    ParBlocks: ArraySize + Div<U2>,
28    Quot<ParBlocks, U2>: ArraySize,
29{
30    const {
31        assert!(matches!(RK, 11 | 13 | 15));
32        assert!(ParBlocks::USIZE % 2 == 0);
33    }
34
35    let mut blocks2 = batch_load(blocks.get_in());
36
37    for block2 in &mut blocks2 {
38        *block2 = _mm256_xor_si256(*block2, keys[0]);
39    }
40    for key in &keys[1..RK - 1] {
41        for block2 in &mut blocks2 {
42            *block2 = _mm256_aesenc_epi128(*block2, *key);
43        }
44    }
45    for block2 in &mut blocks2 {
46        *block2 = _mm256_aesenclast_epi128(*block2, keys[RK - 1]);
47    }
48
49    batch_store(blocks.get_out(), blocks2);
50}
51
52#[inline]
53#[target_feature(enable = "vaes")]
54pub(super) fn batch_decrypt<const RK: usize, ParBlocks>(
55    keys: &RoundKeys2<RK>,
56    mut blocks: InOut<'_, '_, Array<Block, ParBlocks>>,
57) where
58    ParBlocks: ArraySize + Div<U2>,
59    Quot<ParBlocks, U2>: ArraySize,
60{
61    const {
62        assert!(matches!(RK, 11 | 13 | 15));
63        assert!(ParBlocks::USIZE % 2 == 0);
64    }
65
66    let mut blocks2 = batch_load(blocks.get_in());
67
68    for block2 in &mut blocks2 {
69        *block2 = _mm256_xor_si256(*block2, keys[0]);
70    }
71    for key in &keys[1..RK - 1] {
72        for block2 in &mut blocks2 {
73            *block2 = _mm256_aesdec_epi128(*block2, *key);
74        }
75    }
76    for block2 in &mut blocks2 {
77        *block2 = _mm256_aesdeclast_epi128(*block2, keys[RK - 1]);
78    }
79
80    batch_store(blocks.get_out(), blocks2);
81}
82
83#[inline]
84#[target_feature(enable = "avx")]
85fn batch_load<ParBlocks>(blocks: &Array<Block, ParBlocks>) -> BatchBlocks<ParBlocks>
86where
87    ParBlocks: ArraySize + Div<U2>,
88    Quot<ParBlocks, U2>: ArraySize,
89{
90    const { assert!(ParBlocks::USIZE % 2 == 0) }
91
92    let in_ptr: *const __m256i = blocks.as_ptr().cast();
93    // SAFETY: we use unaligned load instruction
94    Array::from_fn(|i| unsafe { _mm256_loadu_si256(in_ptr.add(i)) })
95}
96
97#[inline]
98#[target_feature(enable = "avx")]
99fn batch_store<ParBlocks>(dst: &mut Array<Block, ParBlocks>, blocks: BatchBlocks<ParBlocks>)
100where
101    ParBlocks: ArraySize + Div<U2>,
102    Quot<ParBlocks, U2>: ArraySize,
103{
104    const { assert!(ParBlocks::USIZE % 2 == 0) }
105
106    let dst_ptr: *mut __m256i = dst.as_mut_ptr().cast();
107    for (i, block) in blocks.into_iter().enumerate() {
108        // SAFETY: we use unaligned store instruction
109        unsafe { _mm256_storeu_si256(dst_ptr.add(i), block) }
110    }
111}