native 0.0.1
Vectors, masks and wide register packs for C++26
Loading...
Searching...
No Matches
vaes.h
1// SPDX-License-Identifier: BSD-2-Clause OR Apache-2.0
2#pragma once
3
4#include "native/config.h"
5#include "native/attributes.h"
6#include "native/isa.h"
7#if NATIVE_HOST_X86
8#include <immintrin.h>
9
10namespace native::detail::x86_vaes {
11 template<isa<x86> Arch> requires(Arch.has(x86_feature::aes) && Arch.has(x86_feature::avx))
13 __m128i vaesenc(__m128i state, __m128i key) noexcept {
14 return _mm_aesenc_si128(state, key);
15 }
16
17 template<isa<x86> Arch> requires(Arch.has(x86_feature::vaes) && Arch.has(x86_feature::avx))
19 __m256i vaesenc(__m256i state, __m256i key) noexcept {
20 return _mm256_aesenc_epi128(state, key);
21 }
22
23 template<isa<x86> Arch> requires(Arch.has(x86_feature::vaes) && Arch.has(x86_feature::avx512f))
25 __m512i vaesenc(__m512i state, __m512i key) noexcept {
26 return _mm512_aesenc_epi128(state, key);
27 }
28
29 template<isa<x86> Arch, class... Args>
30 void vaesenc(Args...) = delete;
31
32 template<isa<x86> Arch> requires(Arch.has(x86_feature::aes) && Arch.has(x86_feature::avx))
34 __m128i vaesenclast(__m128i state, __m128i key) noexcept {
35 return _mm_aesenclast_si128(state, key);
36 }
37
38 template<isa<x86> Arch> requires(Arch.has(x86_feature::vaes) && Arch.has(x86_feature::avx))
40 __m256i vaesenclast(__m256i state, __m256i key) noexcept {
41 return _mm256_aesenclast_epi128(state, key);
42 }
43
44 template<isa<x86> Arch> requires(Arch.has(x86_feature::vaes) && Arch.has(x86_feature::avx512f))
46 __m512i vaesenclast(__m512i state, __m512i key) noexcept {
47 return _mm512_aesenclast_epi128(state, key);
48 }
49
50 template<isa<x86> Arch, class... Args>
51 void vaesenclast(Args...) = delete;
52
53 template<isa<x86> Arch> requires(Arch.has(x86_feature::aes) && Arch.has(x86_feature::avx))
55 __m128i vaesdec(__m128i state, __m128i key) noexcept {
56 return _mm_aesdec_si128(state, key);
57 }
58
59 template<isa<x86> Arch> requires(Arch.has(x86_feature::vaes) && Arch.has(x86_feature::avx))
61 __m256i vaesdec(__m256i state, __m256i key) noexcept {
62 return _mm256_aesdec_epi128(state, key);
63 }
64
65 template<isa<x86> Arch> requires(Arch.has(x86_feature::vaes) && Arch.has(x86_feature::avx512f))
67 __m512i vaesdec(__m512i state, __m512i key) noexcept {
68 return _mm512_aesdec_epi128(state, key);
69 }
70
71 template<isa<x86> Arch, class... Args>
72 void vaesdec(Args...) = delete;
73
74 template<isa<x86> Arch> requires(Arch.has(x86_feature::aes) && Arch.has(x86_feature::avx))
76 __m128i vaesdeclast(__m128i state, __m128i key) noexcept {
77 return _mm_aesdeclast_si128(state, key);
78 }
79
80 template<isa<x86> Arch> requires(Arch.has(x86_feature::vaes) && Arch.has(x86_feature::avx))
82 __m256i vaesdeclast(__m256i state, __m256i key) noexcept {
83 return _mm256_aesdeclast_epi128(state, key);
84 }
85
86 template<isa<x86> Arch> requires(Arch.has(x86_feature::vaes) && Arch.has(x86_feature::avx512f))
88 __m512i vaesdeclast(__m512i state, __m512i key) noexcept {
89 return _mm512_aesdeclast_epi128(state, key);
90 }
91
92 template<isa<x86> Arch, class... Args>
93 void vaesdeclast(Args...) = delete;
94
95}
96#endif
97
98#include "native/x86/aes_constant.h"
99
100namespace native::detail::x86_vaes_constant {
101 // Every 128-bit lane is a separate AES state with its own round key. Reuse
102 // the AES substitution and field-mixing primitives without cross-lane work.
103 template<bool Inverse, bool Last, class V>
104 constexpr V round(V state, V key) noexcept {
105 auto input = x86_aes_constant::lanes(state);
106 auto round_key = x86_aes_constant::lanes(key);
107 auto result = input;
108 for (unsigned base = 0; base < V::lanes; base += 16) {
109 std::array<std::uint8_t, 16> block{};
110 for (unsigned column = 0; column < 4; ++column) {
111 for (unsigned row = 0; row < 4; ++row) {
112 auto source = 4 * ((column + (Inverse ? 4 - row : row)) % 4) + row;
113 block[4 * column + row] = x86_aes_constant::substitute<Inverse>(input[base + source]);
114 }
115 }
116 if constexpr (!Last) {
117 block = x86_aes_constant::mix<Inverse>(block);
118 }
119 for (unsigned i = 0; i < 16; ++i) {
120 result[base + i] = block[i] ^ round_key[base + i];
121 }
122 }
123 return V::load(result.data());
124 }
125}
Compiler attributes for host code, with shader-safe shared modifiers.
#define native_inline
inline [[always_inline]]
Definition attributes.h:212
#define native_nodiscard
C++17 [[nodiscard]].
Definition attributes.h:189
#define native_const
[[const]] is not const
Definition attributes.h:108
#define native_target(x)
this indicates a required feature set for the current multiversioned function.
Definition attributes.h:476