Skip to main content

ab_riscv_primitives/instructions/rv64/zk/zkn/
zknd.rs

1//! RV64 Zknd extension
2
3#[cfg(test)]
4mod tests;
5
6use crate::instructions::Instruction;
7use crate::registers::general_purpose::Register;
8use ab_riscv_macros::instruction;
9use core::{fmt, mem};
10
11/// AES key schedule round constant number
12#[derive(Debug, Clone, Copy, PartialEq, Eq)]
13#[repr(u8)]
14pub enum Rv64ZkndKsRnum {
15    R0 = 0x0,
16    R1 = 0x1,
17    R2 = 0x2,
18    R3 = 0x3,
19    R4 = 0x4,
20    R5 = 0x5,
21    R6 = 0x6,
22    R7 = 0x7,
23    R8 = 0x8,
24    R9 = 0x9,
25    Final = 0xA,
26}
27
28const impl From<Rv64ZkndKsRnum> for u8 {
29    #[inline(always)]
30    fn from(rnum: Rv64ZkndKsRnum) -> Self {
31        rnum as u8
32    }
33}
34
35const impl From<Rv64ZkndKsRnum> for usize {
36    #[inline(always)]
37    fn from(rnum: Rv64ZkndKsRnum) -> Self {
38        usize::from(rnum as u8)
39    }
40}
41
42impl fmt::Display for Rv64ZkndKsRnum {
43    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
44        fmt::Display::fmt(&(*self as u8), f)
45    }
46}
47
48impl Rv64ZkndKsRnum {
49    /// Round constants `RC[0..=9]`, indexed by rnum (0-based).
50    /// `RC[rnum]` corresponds to FIPS 197 `Rcon[rnum+1]`.
51    const RCON: [u8; 10] = [0x01, 0x02, 0x04, 0x08, 0x10, 0x20, 0x40, 0x80, 0x1b, 0x36];
52
53    /// Create from raw bits
54    #[inline(always)]
55    pub const fn from_bits(bits: u8) -> Option<Self> {
56        if bits <= Rv64ZkndKsRnum::Final as u8 {
57            // SAFETY: The transmute is safe because `Rv64ZkndKsRnum` is `#[repr(u8)]` enum with
58            // known valid values
59            Some(unsafe { mem::transmute::<u8, Self>(bits) })
60        } else {
61            None
62        }
63    }
64
65    /// Round constant (unless final)
66    #[inline(always)]
67    #[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
68    pub const fn constant(self) -> Option<u8> {
69        if matches!(self, Rv64ZkndKsRnum::Final) {
70            None
71        } else {
72            Some(Self::RCON[usize::from(self)])
73        }
74    }
75}
76
77/// RISC-V RV64 Zknd instructions (AES decryption and key schedule)
78#[instruction]
79#[derive(Debug, Clone, Copy, PartialEq, Eq)]
80pub enum Rv64ZkndInstruction<Reg> {
81    /// AES final round decryption: InvShiftRows + InvSubBytes, no MixColumns
82    Aes64Ds { rd: Reg, rs1: Reg, rs2: Reg },
83    /// AES middle round decryption: InvShiftRows + InvSubBytes + InvMixColumns
84    Aes64Dsm { rd: Reg, rs1: Reg, rs2: Reg },
85    /// AES inverse MixColumns on each 32-bit word of rs1
86    Aes64Im { rd: Reg, rs1: Reg },
87    /// AES key schedule step 1 (rnum in 0..=10)
88    Aes64Ks1i {
89        rd: Reg,
90        rs1: Reg,
91        rnum: Rv64ZkndKsRnum,
92    },
93    /// AES key schedule step 2
94    Aes64Ks2 { rd: Reg, rs1: Reg, rs2: Reg },
95}
96
97#[instruction]
98const impl<Reg> Instruction for Rv64ZkndInstruction<Reg>
99where
100    Reg: [const] Register<Type = u64>,
101{
102    type Reg = Reg;
103
104    #[inline(always)]
105    #[cfg_attr(feature = "no-panic", no_panic_const::no_panic(const))]
106    fn try_decode(instruction: u32) -> Option<Self> {
107        let opcode = (instruction & 0b111_1111) as u8;
108        let rd_bits = ((instruction >> 7) & 0x1f) as u8;
109        let funct3 = ((instruction >> 12) & 0b111) as u8;
110        let rs1_bits = ((instruction >> 15) & 0x1f) as u8;
111        let rs2_bits = ((instruction >> 20) & 0x1f) as u8;
112        let funct7 = ((instruction >> 25) & 0b111_1111) as u8;
113
114        match opcode {
115            // R-type: OP opcode (0x33)
116            //   aes64ds:  funct7=0b001_1101, funct3=0 -> MATCH=0x3a00_0033
117            //   aes64dsm: funct7=0b001_1111, funct3=0 -> MATCH=0x3e00_0033
118            //   aes64ks2: funct7=0b011_1111, funct3=0 -> MATCH=0x7e00_0033
119            0b011_0011 => {
120                if funct3 != 0b000 {
121                    None?;
122                }
123                let rd = Reg::from_bits(rd_bits)?;
124                let rs1 = Reg::from_bits(rs1_bits)?;
125                let rs2 = Reg::from_bits(rs2_bits)?;
126                match funct7 {
127                    0b001_1101 => Some(Self::Aes64Ds { rd, rs1, rs2 }),
128                    0b001_1111 => Some(Self::Aes64Dsm { rd, rs1, rs2 }),
129                    0b011_1111 => Some(Self::Aes64Ks2 { rd, rs1, rs2 }),
130                    _ => None,
131                }
132            }
133            // I-type: OP-IMM opcode (0x13), funct3=0b001
134            //   aes64im:   imm[11:0]=0x300  (funct7=0b001_1000, rs2=0b0_0000) -> MATCH=0x3000_1013
135            //   aes64ks1i: imm[11:5]=0b001_1000, imm[4]=1, imm[3:0]=rnum     -> MATCH=0x3100_1013+
136            0b001_0011 => {
137                if funct3 != 0b001 {
138                    None?;
139                }
140                let rd = Reg::from_bits(rd_bits)?;
141                let rs1 = Reg::from_bits(rs1_bits)?;
142                let imm12 = instruction >> 20;
143                if imm12 == 0x300 {
144                    Some(Self::Aes64Im { rd, rs1 })
145                } else if (imm12 >> 5) == 0b001_1000 && (imm12 & 0b1_0000) != 0 {
146                    // bits[11:5]=0b0011000, bit[4]=1, bits[3:0]=rnum
147                    let rnum = Rv64ZkndKsRnum::from_bits((imm12 & 0xf) as u8)?;
148                    Some(Self::Aes64Ks1i { rd, rs1, rnum })
149                } else {
150                    None
151                }
152            }
153            _ => None,
154        }
155    }
156
157    #[inline(always)]
158    fn alignment() -> u8 {
159        align_of::<u32>() as u8
160    }
161
162    #[inline(always)]
163    fn size(&self) -> u8 {
164        size_of::<u32>() as u8
165    }
166}
167
168#[instruction]
169impl<Reg> fmt::Display for Rv64ZkndInstruction<Reg>
170where
171    Reg: fmt::Display,
172{
173    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
174        match self {
175            Self::Aes64Ds { rd, rs1, rs2 } => write!(f, "aes64ds {rd}, {rs1}, {rs2}"),
176            Self::Aes64Dsm { rd, rs1, rs2 } => write!(f, "aes64dsm {rd}, {rs1}, {rs2}"),
177            Self::Aes64Im { rd, rs1 } => write!(f, "aes64im {rd}, {rs1}"),
178            Self::Aes64Ks1i { rd, rs1, rnum } => write!(f, "aes64ks1i {rd}, {rs1}, {rnum}"),
179            Self::Aes64Ks2 { rd, rs1, rs2 } => write!(f, "aes64ks2 {rd}, {rs1}, {rs2}"),
180        }
181    }
182}