aes/backends/x86_vaes512/
encdec.rs1use super::RoundKeys;
2use crate::Block;
3use cipher::{Array, array::ArraySize, consts::U4, 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 RoundKeys4<const ROUNDS: usize> = [__m512i; ROUNDS];
12
13type BatchBlocks<ParBlocks> = Array<__m512i, Quot<ParBlocks, U4>>;
14
15#[inline]
16#[target_feature(enable = "avx512f")]
17pub(super) fn broadcast_keys<const RK: usize>(keys: &RoundKeys<RK>) -> RoundKeys4<RK> {
18 keys.map(|key| _mm512_broadcast_i32x4(key))
19}
20
21#[inline]
22#[target_feature(enable = "avx512f,vaes")]
23pub(super) fn batch_encrypt<const RK: usize, ParBlocks>(
24 keys: &RoundKeys4<RK>,
25 mut blocks: InOut<'_, '_, Array<Block, ParBlocks>>,
26) where
27 ParBlocks: ArraySize + Div<U4>,
28 Quot<ParBlocks, U4>: ArraySize,
29{
30 const {
31 assert!(matches!(RK, 11 | 13 | 15));
32 assert!(ParBlocks::USIZE % 4 == 0);
33 }
34
35 let mut blocks4 = load(blocks.get_in());
36
37 for block4 in &mut blocks4 {
38 *block4 = _mm512_xor_si512(*block4, keys[0]);
39 }
40 for key in &keys[1..RK - 1] {
41 for block4 in &mut blocks4 {
42 *block4 = _mm512_aesenc_epi128(*block4, *key);
43 }
44 }
45 for block4 in &mut blocks4 {
46 *block4 = _mm512_aesenclast_epi128(*block4, keys[RK - 1]);
47 }
48
49 store(blocks.get_out(), blocks4);
50}
51
52#[inline]
53#[target_feature(enable = "avx512f,vaes")]
54pub(super) fn batch_decrypt<const RK: usize, ParBlocks>(
55 keys: &RoundKeys4<RK>,
56 mut blocks: InOut<'_, '_, Array<Block, ParBlocks>>,
57) where
58 ParBlocks: ArraySize + Div<U4>,
59 Quot<ParBlocks, U4>: ArraySize,
60{
61 const {
62 assert!(matches!(RK, 11 | 13 | 15));
63 assert!(ParBlocks::USIZE % 4 == 0);
64 }
65
66 let mut blocks4 = load(blocks.get_in());
67
68 for block4 in &mut blocks4 {
69 *block4 = _mm512_xor_si512(*block4, keys[0]);
70 }
71 for key in &keys[1..RK - 1] {
72 for block4 in &mut blocks4 {
73 *block4 = _mm512_aesdec_epi128(*block4, *key);
74 }
75 }
76 for block4 in &mut blocks4 {
77 *block4 = _mm512_aesdeclast_epi128(*block4, keys[RK - 1]);
78 }
79
80 store(blocks.get_out(), blocks4);
81}
82
83#[inline]
84#[target_feature(enable = "avx512f")]
85fn load<ParBlocks>(blocks: &Array<Block, ParBlocks>) -> BatchBlocks<ParBlocks>
86where
87 ParBlocks: ArraySize + Div<U4>,
88 Quot<ParBlocks, U4>: ArraySize,
89{
90 const { assert!(ParBlocks::USIZE % 4 == 0) }
91
92 let src_ptr: *const __m512i = blocks.as_ptr().cast();
93 Array::from_fn(|i| unsafe { _mm512_loadu_si512(src_ptr.add(i)) })
95}
96
97#[inline]
98#[target_feature(enable = "avx512f")]
99fn store<ParBlocks>(dst: &mut Array<Block, ParBlocks>, blocks: BatchBlocks<ParBlocks>)
100where
101 ParBlocks: ArraySize + Div<U4>,
102 Quot<ParBlocks, U4>: ArraySize,
103{
104 const { assert!(ParBlocks::USIZE % 4 == 0) }
105
106 let dst_ptr: *mut __m512i = dst.as_mut_ptr().cast();
107 for (i, block) in blocks.into_iter().enumerate() {
108 unsafe { _mm512_storeu_si512(dst_ptr.add(i), block) }
110 }
111}