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