mod aes;
mod ghash;
mod sha2;
mod sm3;
mod sm4;
use crate::exec::compute::vector::regfile::VectorRegFile;
use crate::isa::op::CryptoOp;
use crate::isa::rvv::{ElemIdx, Sew, VRegIdx};
use aes::{
aes_kf1, aes_kf2, aes_round_dec, aes_round_dec_final, aes_round_enc, aes_round_enc_final,
aes_round_zero,
};
use ghash::gf128_mul;
use sha2::{sha256_compress_high, sha256_compress_low, sha256_ms};
use sm3::{sm3_c, sm3_me};
use sm4::{sm4_k, sm4_r};
const EGS_AES: usize = 4;
const EGS_SM3: usize = 8;
const EGS_SHA: usize = 4;
#[inline]
fn read_egs4_u32(vpr: &impl VectorRegFile, vreg: VRegIdx, base_elem: usize) -> [u32; 4] {
let mut out = [0u32; 4];
for (i, slot) in out.iter_mut().enumerate() {
*slot = vpr.read_element(vreg, ElemIdx::new(base_elem + i), Sew::E32) as u32;
}
out
}
#[inline]
fn write_egs4_u32(vpr: &mut impl VectorRegFile, vreg: VRegIdx, base_elem: usize, vals: [u32; 4]) {
for (i, v) in vals.iter().enumerate() {
vpr.write_element(vreg, ElemIdx::new(base_elem + i), Sew::E32, u64::from(*v));
}
}
#[inline]
fn read_egs8_u32(vpr: &impl VectorRegFile, vreg: VRegIdx, base_elem: usize) -> [u32; 8] {
let mut out = [0u32; 8];
for (i, slot) in out.iter_mut().enumerate() {
*slot = vpr.read_element(vreg, ElemIdx::new(base_elem + i), Sew::E32) as u32;
}
out
}
#[inline]
fn write_egs8_u32(vpr: &mut impl VectorRegFile, vreg: VRegIdx, base_elem: usize, vals: [u32; 8]) {
for (i, v) in vals.iter().enumerate() {
vpr.write_element(vreg, ElemIdx::new(base_elem + i), Sew::E32, u64::from(*v));
}
}
#[allow(clippy::too_many_arguments)]
pub fn execute_crypto(
op: CryptoOp,
vpr: &mut impl VectorRegFile,
vd_idx: VRegIdx,
vs2_idx: VRegIdx,
vs1_idx: VRegIdx,
vstart: usize,
vl: usize,
inst: u32,
broadcast_vs2: bool,
) {
let egs = if matches!(op, CryptoOp::Sm3Me | CryptoOp::Sm3C) { EGS_SM3 } else { EGS_AES };
let _ = EGS_SHA;
if vl < egs {
return;
}
let key_base = |base: usize| if broadcast_vs2 { 0 } else { base };
let mut base = (vstart / egs) * egs;
while base + egs <= vl {
match op {
CryptoOp::AesEm => {
let state = read_egs4_u32(vpr, vd_idx, base);
let key = read_egs4_u32(vpr, vs2_idx, key_base(base));
let r = aes_round_enc(state, key);
write_egs4_u32(vpr, vd_idx, base, r);
}
CryptoOp::AesEf => {
let state = read_egs4_u32(vpr, vd_idx, base);
let key = read_egs4_u32(vpr, vs2_idx, key_base(base));
let r = aes_round_enc_final(state, key);
write_egs4_u32(vpr, vd_idx, base, r);
}
CryptoOp::AesDm => {
let state = read_egs4_u32(vpr, vd_idx, base);
let key = read_egs4_u32(vpr, vs2_idx, key_base(base));
let r = aes_round_dec(state, key);
write_egs4_u32(vpr, vd_idx, base, r);
}
CryptoOp::AesDf => {
let state = read_egs4_u32(vpr, vd_idx, base);
let key = read_egs4_u32(vpr, vs2_idx, key_base(base));
let r = aes_round_dec_final(state, key);
write_egs4_u32(vpr, vd_idx, base, r);
}
CryptoOp::AesZ => {
let state = read_egs4_u32(vpr, vd_idx, base);
let key = read_egs4_u32(vpr, vs2_idx, 0);
let r = aes_round_zero(state, key);
write_egs4_u32(vpr, vd_idx, base, r);
}
CryptoOp::AesKf1 => {
let rnd = (inst >> 15) & 0x1f;
let prev = read_egs4_u32(vpr, vs2_idx, base);
let r = aes_kf1(prev, rnd);
write_egs4_u32(vpr, vd_idx, base, r);
}
CryptoOp::AesKf2 => {
let rnd = (inst >> 15) & 0x1f;
let curr = read_egs4_u32(vpr, vd_idx, base);
let prev = read_egs4_u32(vpr, vs2_idx, base);
let r = aes_kf2(curr, prev, rnd);
write_egs4_u32(vpr, vd_idx, base, r);
}
CryptoOp::Sha2Ms => {
let vd = read_egs4_u32(vpr, vd_idx, base);
let vs2 = read_egs4_u32(vpr, vs2_idx, base);
let vs1 = read_egs4_u32(vpr, vs1_idx, base);
let r = sha256_ms(vd, vs2, vs1);
write_egs4_u32(vpr, vd_idx, base, r);
}
CryptoOp::Sha2Cl => {
let vd = read_egs4_u32(vpr, vd_idx, base);
let vs2 = read_egs4_u32(vpr, vs2_idx, base);
let vs1 = read_egs4_u32(vpr, vs1_idx, base);
let r = sha256_compress_low(vd, vs2, vs1);
write_egs4_u32(vpr, vd_idx, base, r);
}
CryptoOp::Sha2Ch => {
let vd = read_egs4_u32(vpr, vd_idx, base);
let vs2 = read_egs4_u32(vpr, vs2_idx, base);
let vs1 = read_egs4_u32(vpr, vs1_idx, base);
let r = sha256_compress_high(vd, vs2, vs1);
write_egs4_u32(vpr, vd_idx, base, r);
}
CryptoOp::Sm3Me => {
let vs1 = read_egs8_u32(vpr, vs1_idx, base);
let vs2 = read_egs8_u32(vpr, vs2_idx, base);
let r = sm3_me(vs1, vs2);
write_egs8_u32(vpr, vd_idx, base, r);
}
CryptoOp::Sm3C => {
let rnd = (inst >> 15) & 0x1f;
let vd = read_egs8_u32(vpr, vd_idx, base);
let vs2 = read_egs8_u32(vpr, vs2_idx, base);
let r = sm3_c(vd, vs2, rnd);
write_egs8_u32(vpr, vd_idx, base, r);
}
CryptoOp::Sm4R => {
let vd = read_egs4_u32(vpr, vd_idx, base);
let vs2 = read_egs4_u32(vpr, vs2_idx, key_base(base));
let r = sm4_r(vd, vs2);
write_egs4_u32(vpr, vd_idx, base, r);
}
CryptoOp::Sm4K => {
let rnd = (inst >> 15) & 0x1f;
let prev = read_egs4_u32(vpr, vs2_idx, base);
let r = sm4_k(prev, rnd);
write_egs4_u32(vpr, vd_idx, base, r);
}
CryptoOp::Gmul => {
let vd = read_egs4_u32(vpr, vd_idx, base);
let vs2 = read_egs4_u32(vpr, vs2_idx, base);
let r = gf128_mul(vd, vs2);
write_egs4_u32(vpr, vd_idx, base, r);
}
CryptoOp::Ghsh => {
let vd = read_egs4_u32(vpr, vd_idx, base);
let vs2 = read_egs4_u32(vpr, vs2_idx, base);
let vs1 = read_egs4_u32(vpr, vs1_idx, base);
let xored = [vd[0] ^ vs1[0], vd[1] ^ vs1[1], vd[2] ^ vs1[2], vd[3] ^ vs1[3]];
let r = gf128_mul(xored, vs2);
write_egs4_u32(vpr, vd_idx, base, r);
}
}
base += egs;
}
}
#[cfg(test)]
mod tests {
use super::*;
use aes::aes_kf1;
use aes::aes_round_enc;
fn bytes_to_words(b: [u8; 16]) -> [u32; 4] {
[
u32::from_le_bytes([b[0], b[1], b[2], b[3]]),
u32::from_le_bytes([b[4], b[5], b[6], b[7]]),
u32::from_le_bytes([b[8], b[9], b[10], b[11]]),
u32::from_le_bytes([b[12], b[13], b[14], b[15]]),
]
}
#[test]
fn aes_round_enc_matches_fips197_round1() {
let state_in = bytes_to_words([
0x19, 0x3d, 0xe3, 0xbe, 0xa0, 0xf4, 0xe2, 0x2b, 0x9a, 0xc6, 0x8d, 0x2a, 0xe9, 0xf8,
0x48, 0x08,
]);
let round_key = bytes_to_words([
0xa0, 0xfa, 0xfe, 0x17, 0x88, 0x54, 0x2c, 0xb1, 0x23, 0xa3, 0x39, 0x39, 0x2a, 0x6c,
0x76, 0x05,
]);
let expected = bytes_to_words([
0xa4, 0x9c, 0x7f, 0xf2, 0x68, 0x9f, 0x35, 0x2b, 0x6b, 0x5b, 0xea, 0x43, 0x02, 0x6a,
0x50, 0x49,
]);
assert_eq!(aes_round_enc(state_in, round_key), expected);
}
#[test]
fn aes_kf1_matches_fips197_round1() {
let key0 = bytes_to_words([
0x2b, 0x7e, 0x15, 0x16, 0x28, 0xae, 0xd2, 0xa6, 0xab, 0xf7, 0x15, 0x88, 0x09, 0xcf,
0x4f, 0x3c,
]);
let expected = bytes_to_words([
0xa0, 0xfa, 0xfe, 0x17, 0x88, 0x54, 0x2c, 0xb1, 0x23, 0xa3, 0x39, 0x39, 0x2a, 0x6c,
0x76, 0x05,
]);
assert_eq!(aes_kf1(key0, 1), expected);
}
}