4#include "native/x86/integer_constant.h"
8namespace native::detail::x86_aes_constant {
9 using x86_instruction_constant::field_inverse;
10 using x86_instruction_constant::field_product;
11 using x86_instruction_constant::lanes;
13 template<
bool Inverse>
14 constexpr std::uint8_t substitute(std::uint8_t value)
noexcept {
15 if constexpr (Inverse) {
16 return field_inverse(
static_cast<std::uint8_t
>(
17 std::rotl(value, 1) ^ std::rotl(value, 3) ^ std::rotl(value, 6) ^ 5));
19 auto x = field_inverse(value);
20 return static_cast<std::uint8_t
>(x ^ std::rotl(x, 1) ^ std::rotl(x, 2) ^
21 std::rotl(x, 3) ^ std::rotl(x, 4) ^ 0x63);
25 template<
bool Inverse>
26 constexpr auto mix(std::array<std::uint8_t, 16> input)
noexcept {
27 std::array<std::uint8_t, 16> result{};
28 constexpr std::uint8_t coefficients[2][4]{{2, 3, 1, 1}, {14, 11, 13, 9}};
29 for (
unsigned column = 0; column < 4; ++column) {
30 for (
unsigned row = 0; row < 4; ++row) {
31 for (
unsigned term = 0; term < 4; ++term) {
32 result[4 * column + row] ^= field_product(input[4 * column + term],
33 coefficients[Inverse][(term + 4 - row) % 4]);
40 template<
bool Inverse,
bool Last,
class V>
41 constexpr V round(V state, V key)
noexcept {
42 auto input = lanes(state);
44 auto round_key = lanes(key);
45 for (
unsigned column = 0; column < 4; ++column) {
46 for (
unsigned row = 0; row < 4; ++row) {
47 auto source = 4 * ((column + (Inverse ? 4 - row : row)) % 4) + row;
48 result[4 * column + row] = substitute<Inverse>(input[source]);
51 if constexpr (!Last) {
52 result = mix<Inverse>(result);
54 for (
unsigned i = 0; i < 16; ++i) {
55 result[i] ^= round_key[i];
57 return V::load(result.data());
61 constexpr V inverse_mix(V state)
noexcept {
62 auto result = mix<true>(lanes(state));
63 return V::load(result.data());
66 template<
unsigned Imm8,
class V>
67 constexpr V keygen(V state)
noexcept {
68 auto input = lanes(state);
70 for (
unsigned half = 0; half < 2; ++half) {
71 auto source = 8 * half + 4;
72 for (
unsigned i = 0; i < 4; ++i) {
73 result[8 * half + i] = substitute<false>(input[source + i]);
74 result[8 * half + 4 + i] = substitute<false>(input[source + (i + 1) % 4]);
76 result[8 * half + 4] ^= Imm8;
78 return V::load(result.data());