Merge pull request #2391 from lioncash/aes

IR: Allow specifying register size for AES enc/dec ops and PCLMUL
This commit is contained in:
Ryan Houdek authored and GitHub committed 2023-02-08 17:11:03 -08:00
commit 8689038533
4 files changed
+139 -49

No files matched your search

@@ -17,37 +17,73 @@ DEF_OP(AESImc) {
}
DEF_OP(AESEnc) {
auto Op = IROp->C<IR::IROp_VAESEnc>();
const auto Op = IROp->C<IR::IROp_VAESEnc>();
const auto OpSize = IROp->Size;
const auto Dst = GetVReg(Node);
const auto Key = GetVReg(Op->Key.ID());
const auto State = GetVReg(Op->State.ID());
LOGMAN_THROW_AA_FMT(OpSize == Core::CPUState::XMM_SSE_REG_SIZE,
"Currently only supports 128-bit operations.");
eor(VTMP2.Q(), VTMP2.Q(), VTMP2.Q());
mov(VTMP1.Q(), GetVReg(Op->State.ID()).Q());
mov(VTMP1.Q(), State.Q());
aese(VTMP1, VTMP2);
aesmc(VTMP1, VTMP1);
eor(GetVReg(Node).Q(), VTMP1.Q(), GetVReg(Op->Key.ID()).Q());
eor(Dst.Q(), VTMP1.Q(), Key.Q());
}
DEF_OP(AESEncLast) {
auto Op = IROp->C<IR::IROp_VAESEncLast>();
const auto Op = IROp->C<IR::IROp_VAESEncLast>();
const auto OpSize = IROp->Size;
const auto Dst = GetVReg(Node);
const auto Key = GetVReg(Op->Key.ID());
const auto State = GetVReg(Op->State.ID());
LOGMAN_THROW_AA_FMT(OpSize == Core::CPUState::XMM_SSE_REG_SIZE,
"Currently only supports 128-bit operations.");
eor(VTMP2.Q(), VTMP2.Q(), VTMP2.Q());
mov(VTMP1.Q(), GetVReg(Op->State.ID()).Q());
mov(VTMP1.Q(), State.Q());
aese(VTMP1, VTMP2);
eor(GetVReg(Node).Q(), VTMP1.Q(), GetVReg(Op->Key.ID()).Q());
eor(Dst.Q(), VTMP1.Q(), Key.Q());
}
DEF_OP(AESDec) {
auto Op = IROp->C<IR::IROp_VAESDec>();
const auto Op = IROp->C<IR::IROp_VAESDec>();
const auto OpSize = IROp->Size;
const auto Dst = GetVReg(Node);
const auto Key = GetVReg(Op->Key.ID());
const auto State = GetVReg(Op->State.ID());
LOGMAN_THROW_AA_FMT(OpSize == Core::CPUState::XMM_SSE_REG_SIZE,
"Currently only supports 128-bit operations.");
eor(VTMP2.Q(), VTMP2.Q(), VTMP2.Q());
mov(VTMP1.Q(), GetVReg(Op->State.ID()).Q());
mov(VTMP1.Q(), State.Q());
aesd(VTMP1, VTMP2);
aesimc(VTMP1, VTMP1);
eor(GetVReg(Node).Q(), VTMP1.Q(), GetVReg(Op->Key.ID()).Q());
eor(Dst.Q(), VTMP1.Q(), Key.Q());
}
DEF_OP(AESDecLast) {
auto Op = IROp->C<IR::IROp_VAESDecLast>();
const auto Op = IROp->C<IR::IROp_VAESDecLast>();
const auto OpSize = IROp->Size;
const auto Dst = GetVReg(Node);
const auto Key = GetVReg(Op->Key.ID());
const auto State = GetVReg(Op->State.ID());
LOGMAN_THROW_AA_FMT(OpSize == Core::CPUState::XMM_SSE_REG_SIZE,
"Currently only supports 128-bit operations.");
eor(VTMP2.Q(), VTMP2.Q(), VTMP2.Q());
mov(VTMP1.Q(), GetVReg(Op->State.ID()).Q());
mov(VTMP1.Q(), State.Q());
aesd(VTMP1, VTMP2);
eor(GetVReg(Node).Q(), VTMP1.Q(), GetVReg(Op->Key.ID()).Q());
eor(Dst.Q(), VTMP1.Q(), Key.Q());
}
DEF_OP(AESKeyGenAssist) {
@@ -108,11 +144,15 @@ DEF_OP(CRC32) {
}
DEF_OP(PCLMUL) {
auto Op = IROp->C<IR::IROp_PCLMUL>();
const auto Op = IROp->C<IR::IROp_PCLMUL>();
const auto OpSize = IROp->Size;
auto Dst = GetVReg(Node);
auto Src1 = GetVReg(Op->Src1.ID());
auto Src2 = GetVReg(Op->Src2.ID());
const auto Dst = GetVReg(Node);
const auto Src1 = GetVReg(Op->Src1.ID());
const auto Src2 = GetVReg(Op->Src2.ID());
LOGMAN_THROW_AA_FMT(OpSize == Core::CPUState::XMM_SSE_REG_SIZE,
"Currently only supports 128-bit operations.");
switch (Op->Selector) {
case 0b00000000:
@@ -21,23 +21,67 @@ DEF_OP(AESImc) {
}
DEF_OP(AESEnc) {
auto Op = IROp->C<IR::IROp_VAESEnc>();
vaesenc(GetDst(Node), GetSrc(Op->State.ID()), GetSrc(Op->Key.ID()));
const auto Op = IROp->C<IR::IROp_VAESEnc>();
const auto OpSize = IROp->Size;
const auto Is256Bit = OpSize == Core::CPUState::XMM_AVX_REG_SIZE;
const auto Dst = GetDst(Node);
const auto Key = GetSrc(Op->Key.ID());
const auto State = GetSrc(Op->State.ID());
if (Is256Bit) {
vaesenc(ToYMM(Dst), ToYMM(State), ToYMM(Key));
} else {
vaesenc(Dst, State, Key);
}
}
DEF_OP(AESEncLast) {
auto Op = IROp->C<IR::IROp_VAESEncLast>();
vaesenclast(GetDst(Node), GetSrc(Op->State.ID()), GetSrc(Op->Key.ID()));
const auto Op = IROp->C<IR::IROp_VAESEncLast>();
const auto OpSize = IROp->Size;
const auto Is256Bit = OpSize == Core::CPUState::XMM_AVX_REG_SIZE;
const auto Dst = GetDst(Node);
const auto Key = GetSrc(Op->Key.ID());
const auto State = GetSrc(Op->State.ID());
if (Is256Bit) {
vaesenclast(ToYMM(Dst), ToYMM(State), ToYMM(Key));
} else {
vaesenclast(Dst, State, Key);
}
}
DEF_OP(AESDec) {
auto Op = IROp->C<IR::IROp_VAESDec>();
vaesdec(GetDst(Node), GetSrc(Op->State.ID()), GetSrc(Op->Key.ID()));
const auto Op = IROp->C<IR::IROp_VAESDec>();
const auto OpSize = IROp->Size;
const auto Is256Bit = OpSize == Core::CPUState::XMM_AVX_REG_SIZE;
const auto Dst = GetDst(Node);
const auto Key = GetSrc(Op->Key.ID());
const auto State = GetSrc(Op->State.ID());
if (Is256Bit) {
vaesdec(ToYMM(Dst), ToYMM(State), ToYMM(Key));
} else {
vaesdec(Dst, State, Key);
}
}
DEF_OP(AESDecLast) {
auto Op = IROp->C<IR::IROp_VAESDecLast>();
vaesdeclast(GetDst(Node), GetSrc(Op->State.ID()), GetSrc(Op->Key.ID()));
const auto Op = IROp->C<IR::IROp_VAESDecLast>();
const auto OpSize = IROp->Size;
const auto Is256Bit = OpSize == Core::CPUState::XMM_AVX_REG_SIZE;
const auto Dst = GetDst(Node);
const auto Key = GetSrc(Op->Key.ID());
const auto State = GetSrc(Op->State.ID());
if (Is256Bit) {
vaesdeclast(ToYMM(Dst), ToYMM(State), ToYMM(Key));
} else {
vaesdeclast(Dst, State, Key);
}
}
DEF_OP(AESKeyGenAssist) {
@@ -76,18 +120,24 @@ DEF_OP(CRC32) {
}
DEF_OP(PCLMUL) {
auto Op = IROp->C<IR::IROp_PCLMUL>();
const auto Op = IROp->C<IR::IROp_PCLMUL>();
const auto OpSize = IROp->Size;
const auto Is256Bit = OpSize == Core::CPUState::XMM_AVX_REG_SIZE;
auto Dst = GetDst(Node);
auto Src1 = GetSrc(Op->Src1.ID());
auto Src2 = GetSrc(Op->Src2.ID());
const auto Dst = GetDst(Node);
const auto Src1 = GetSrc(Op->Src1.ID());
const auto Src2 = GetSrc(Op->Src2.ID());
switch (Op->Selector) {
case 0b00000000:
case 0b00000001:
case 0b00010000:
case 0b00010001:
vpclmulqdq(Dst, Src1, Src2, Op->Selector);
if (Is256Bit) {
vpclmulqdq(ToYMM(Dst), ToYMM(Src1), ToYMM(Src2), Op->Selector);
} else {
vpclmulqdq(Dst, Src1, Src2, Op->Selector);
}
break;
default:
LOGMAN_MSG_A_FMT("Unknown PCLMUL selector: {}", Op->Selector);
@@ -280,7 +280,7 @@ void OpDispatchBuilder::VAESIMCOp(OpcodeArgs) {
void OpDispatchBuilder::AESEncOp(OpcodeArgs) {
OrderedNode *Dest = LoadSource(FPRClass, Op, Op->Dest, Op->Flags, -1);
OrderedNode *Src = LoadSource(FPRClass, Op, Op->Src[0], Op->Flags, -1);
OrderedNode *Result = _VAESEnc(Dest, Src);
OrderedNode *Result = _VAESEnc(16, Dest, Src);
StoreResult(FPRClass, Op, Result, -1);
}
@@ -293,7 +293,7 @@ void OpDispatchBuilder::VAESEncOp(OpcodeArgs) {
OrderedNode *State = LoadSource(FPRClass, Op, Op->Src[0], Op->Flags, -1);
OrderedNode *Key = LoadSource(FPRClass, Op, Op->Src[1], Op->Flags, -1);
OrderedNode *Result = _VAESEnc(State, Key);
OrderedNode *Result = _VAESEnc(DstSize, State, Key);
if (Is128Bit) {
Result = _VMov(16, Result);
@@ -304,7 +304,7 @@ void OpDispatchBuilder::VAESEncOp(OpcodeArgs) {
void OpDispatchBuilder::AESEncLastOp(OpcodeArgs) {
OrderedNode *Dest = LoadSource(FPRClass, Op, Op->Dest, Op->Flags, -1);
OrderedNode *Src = LoadSource(FPRClass, Op, Op->Src[0], Op->Flags, -1);
OrderedNode *Result = _VAESEncLast(Dest, Src);
OrderedNode *Result = _VAESEncLast(16, Dest, Src);
StoreResult(FPRClass, Op, Result, -1);
}
@@ -317,7 +317,7 @@ void OpDispatchBuilder::VAESEncLastOp(OpcodeArgs) {
OrderedNode *State = LoadSource(FPRClass, Op, Op->Src[0], Op->Flags, -1);
OrderedNode *Key = LoadSource(FPRClass, Op, Op->Src[1], Op->Flags, -1);
OrderedNode *Result = _VAESEncLast(State, Key);
OrderedNode *Result = _VAESEncLast(DstSize, State, Key);
if (Is128Bit) {
Result = _VMov(16, Result);
@@ -328,7 +328,7 @@ void OpDispatchBuilder::VAESEncLastOp(OpcodeArgs) {
void OpDispatchBuilder::AESDecOp(OpcodeArgs) {
OrderedNode *Dest = LoadSource(FPRClass, Op, Op->Dest, Op->Flags, -1);
OrderedNode *Src = LoadSource(FPRClass, Op, Op->Src[0], Op->Flags, -1);
OrderedNode *Result = _VAESDec(Dest, Src);
OrderedNode *Result = _VAESDec(16, Dest, Src);
StoreResult(FPRClass, Op, Result, -1);
}
@@ -341,7 +341,7 @@ void OpDispatchBuilder::VAESDecOp(OpcodeArgs) {
OrderedNode *State = LoadSource(FPRClass, Op, Op->Src[0], Op->Flags, -1);
OrderedNode *Key = LoadSource(FPRClass, Op, Op->Src[1], Op->Flags, -1);
OrderedNode *Result = _VAESDec(State, Key);
OrderedNode *Result = _VAESDec(DstSize, State, Key);
if (Is128Bit) {
Result = _VMov(16, Result);
@@ -352,7 +352,7 @@ void OpDispatchBuilder::VAESDecOp(OpcodeArgs) {
void OpDispatchBuilder::AESDecLastOp(OpcodeArgs) {
OrderedNode *Dest = LoadSource(FPRClass, Op, Op->Dest, Op->Flags, -1);
OrderedNode *Src = LoadSource(FPRClass, Op, Op->Src[0], Op->Flags, -1);
OrderedNode *Result = _VAESDecLast(Dest, Src);
OrderedNode *Result = _VAESDecLast(16, Dest, Src);
StoreResult(FPRClass, Op, Result, -1);
}
@@ -365,7 +365,7 @@ void OpDispatchBuilder::VAESDecLastOp(OpcodeArgs) {
OrderedNode *State = LoadSource(FPRClass, Op, Op->Src[0], Op->Flags, -1);
OrderedNode *Key = LoadSource(FPRClass, Op, Op->Src[1], Op->Flags, -1);
OrderedNode *Result = _VAESDecLast(State, Key);
OrderedNode *Result = _VAESDecLast(DstSize, State, Key);
if (Is128Bit) {
Result = _VMov(16, Result);
@@ -399,7 +399,7 @@ void OpDispatchBuilder::PCLMULQDQOp(OpcodeArgs) {
OrderedNode *Src = LoadSource(FPRClass, Op, Op->Src[0], Op->Flags, -1);
const auto Selector = static_cast<uint8_t>(Op->Src[1].Data.Literal.Value);
auto Res = _PCLMUL(Dest, Src, Selector);
auto Res = _PCLMUL(16, Dest, Src, Selector);
StoreResult(FPRClass, Op, Res, -1);
}
@@ -413,7 +413,7 @@ void OpDispatchBuilder::VPCLMULQDQOp(OpcodeArgs) {
OrderedNode *Src2 = LoadSource(FPRClass, Op, Op->Src[1], Op->Flags, -1);
const auto Selector = static_cast<uint8_t>(Op->Src[2].Data.Literal.Value);
OrderedNode *Res = _PCLMUL(Src1, Src2, Selector);
OrderedNode *Res = _PCLMUL(DstSize, Src1, Src2, Selector);
if (Is128Bit) {
Res = _VMov(16, Res);
}
+10 -10
View File
@@ -1418,21 +1418,21 @@
"Desc": "Does a stage of the inverse mix column transformation",
"DestSize": "16"
},
"FPR = VAESEnc FPR:$State, FPR:$Key": {
"FPR = VAESEnc u8:#RegisterSize, FPR:$State, FPR:$Key": {
"Desc": "Does a step of AES encryption",
"DestSize": "16"
"DestSize": "RegisterSize"
},
"FPR = VAESEncLast FPR:$State, FPR:$Key": {
"FPR = VAESEncLast u8:#RegisterSize, FPR:$State, FPR:$Key": {
"Desc": "Does the last step of AES encryption",
"DestSize": "16"
"DestSize": "RegisterSize"
},
"FPR = VAESDec FPR:$State, FPR:$Key": {
"FPR = VAESDec u8:#RegisterSize, FPR:$State, FPR:$Key": {
"Desc": "Does a step of AES decryption",
"DestSize": "16"
"DestSize": "RegisterSize"
},
"FPR = VAESDecLast FPR:$State, FPR:$Key": {
"FPR = VAESDecLast u8:#RegisterSize, FPR:$State, FPR:$Key": {
"Desc": "Does the last step of AES decryption",
"DestSize": "16"
"DestSize": "RegisterSize"
},
"FPR = VAESKeyGenAssist FPR:$Src, u8:$RCON": {
"Desc": "Assists in key generation",
@@ -1443,7 +1443,7 @@
],
"DestSize": "std::max<uint8_t>(4, GetOpSize(_Src1))"
},
"FPR = PCLMUL FPR:$Src1, FPR:$Src2, u8:$Selector": {
"FPR = PCLMUL u8:#RegisterSize, FPR:$Src1, FPR:$Src2, u8:$Selector": {
"Desc": [
"Performs carryless multiplication of 64-bit elements depending on the selector.",
"Selector = 0b00000000: Uses low 64-bit elements from both input vectors",
@@ -1451,7 +1451,7 @@
"Selector = 0b00010000: Uses low 64-bit element from Src1 and high 64-bit element from Src2",
"Selector = 0b00010001: Uses high 64-bit elements from both input vectors"
],
"DestSize": "16"
"DestSize": "RegisterSize"
}
},
"F64": {