Skip to main content

ic_cipher/aes/
portable.rs

1//! The portable, constant-time AES backend.
2//!
3//! Works on every target, computes its S-box algebraically rather than reading
4//! a table, and is the reference against which the accelerated backends are
5//! differentially tested.
6
7use crate::gf;
8use ic_core::{ensure, Result, Zeroize};
9
10/// AES block size in bytes.
11pub const BLOCK_LEN: usize = 16;
12
13/// Round keys for the widest schedule (AES-256 has 15).
14pub(crate) const MAX_ROUND_KEYS: usize = 15 * BLOCK_LEN;
15
16/// Round-constant sequence for the key schedule, `rcon[i] = x^i` in GF(2^8).
17fn rcon(i: usize) -> u8 {
18    let mut c = 1u8;
19    for _ in 1..i {
20        c = gf::xtime(c);
21    }
22    c
23}
24
25/// An expanded AES key schedule.
26///
27/// This is the *only* key expansion in the library. The accelerated backends
28/// consume its output rather than reimplementing it: expansion happens once per
29/// key and is not on the hot path, so there is nothing to gain from a second
30/// implementation and a subtle divergence to lose.
31///
32/// Dropping the schedule zeroizes it, so round keys never outlive the value.
33#[derive(Clone)]
34pub struct Schedule {
35    round_keys: [u8; MAX_ROUND_KEYS],
36    pub(crate) rounds: usize,
37}
38
39impl Drop for Schedule {
40    fn drop(&mut self) {
41        self.round_keys.zeroize();
42    }
43}
44
45impl Schedule {
46    /// Expand a 16, 24, or 32 byte key.
47    pub fn expand(key: &[u8]) -> Result<Self> {
48        let nk = key.len() / 4;
49        let rounds = match key.len() {
50            16 => 10,
51            24 => 12,
52            32 => 14,
53            _ => {
54                return Err(ic_core::err!(
55                    InvalidLength,
56                    "aes key must be 16, 24, or 32 bytes"
57                ))
58            }
59        };
60        let total_words = (rounds + 1) * 4;
61        let mut rk = [0u8; MAX_ROUND_KEYS];
62        rk[..key.len()].copy_from_slice(key);
63
64        for i in nk..total_words {
65            let mut t = [
66                rk[(i - 1) * 4],
67                rk[(i - 1) * 4 + 1],
68                rk[(i - 1) * 4 + 2],
69                rk[(i - 1) * 4 + 3],
70            ];
71            if i % nk == 0 {
72                t.rotate_left(1);
73                for b in t.iter_mut() {
74                    *b = gf::sbox(*b);
75                }
76                t[0] ^= rcon(i / nk);
77            } else if nk > 6 && i % nk == 4 {
78                for b in t.iter_mut() {
79                    *b = gf::sbox(*b);
80                }
81            }
82            for j in 0..4 {
83                rk[i * 4 + j] = rk[(i - nk) * 4 + j] ^ t[j];
84            }
85        }
86        Ok(Self {
87            round_keys: rk,
88            rounds,
89        })
90    }
91
92    /// The round key for round `round`.
93    #[inline]
94    pub(crate) fn round_key(&self, round: usize) -> &[u8] {
95        &self.round_keys[round * BLOCK_LEN..(round + 1) * BLOCK_LEN]
96    }
97}
98
99#[inline]
100fn add_round_key(state: &mut [u8; BLOCK_LEN], rk: &[u8]) {
101    for i in 0..BLOCK_LEN {
102        state[i] ^= rk[i];
103    }
104}
105
106#[inline]
107fn sub_bytes(state: &mut [u8; BLOCK_LEN]) {
108    for b in state.iter_mut() {
109        *b = gf::sbox(*b);
110    }
111}
112
113#[inline]
114fn inv_sub_bytes(state: &mut [u8; BLOCK_LEN]) {
115    for b in state.iter_mut() {
116        *b = gf::inv_sbox(*b);
117    }
118}
119
120/// ShiftRows on the column-major AES state: row `r` rotates left by `r`.
121#[inline]
122fn shift_rows(s: &mut [u8; BLOCK_LEN]) {
123    let t = *s;
124    for c in 0..4 {
125        for r in 0..4 {
126            s[c * 4 + r] = t[((c + r) % 4) * 4 + r];
127        }
128    }
129}
130
131#[inline]
132fn inv_shift_rows(s: &mut [u8; BLOCK_LEN]) {
133    let t = *s;
134    for c in 0..4 {
135        for r in 0..4 {
136            s[((c + r) % 4) * 4 + r] = t[c * 4 + r];
137        }
138    }
139}
140
141#[inline]
142fn mix_columns(s: &mut [u8; BLOCK_LEN]) {
143    for c in 0..4 {
144        let col = [s[c * 4], s[c * 4 + 1], s[c * 4 + 2], s[c * 4 + 3]];
145        s[c * 4] = gf::xtime(col[0]) ^ (gf::xtime(col[1]) ^ col[1]) ^ col[2] ^ col[3];
146        s[c * 4 + 1] = col[0] ^ gf::xtime(col[1]) ^ (gf::xtime(col[2]) ^ col[2]) ^ col[3];
147        s[c * 4 + 2] = col[0] ^ col[1] ^ gf::xtime(col[2]) ^ (gf::xtime(col[3]) ^ col[3]);
148        s[c * 4 + 3] = (gf::xtime(col[0]) ^ col[0]) ^ col[1] ^ col[2] ^ gf::xtime(col[3]);
149    }
150}
151
152#[inline]
153fn inv_mix_columns(s: &mut [u8; BLOCK_LEN]) {
154    for c in 0..4 {
155        let col = [s[c * 4], s[c * 4 + 1], s[c * 4 + 2], s[c * 4 + 3]];
156        s[c * 4] =
157            gf::mul(col[0], 14) ^ gf::mul(col[1], 11) ^ gf::mul(col[2], 13) ^ gf::mul(col[3], 9);
158        s[c * 4 + 1] =
159            gf::mul(col[0], 9) ^ gf::mul(col[1], 14) ^ gf::mul(col[2], 11) ^ gf::mul(col[3], 13);
160        s[c * 4 + 2] =
161            gf::mul(col[0], 13) ^ gf::mul(col[1], 9) ^ gf::mul(col[2], 14) ^ gf::mul(col[3], 11);
162        s[c * 4 + 3] =
163            gf::mul(col[0], 11) ^ gf::mul(col[1], 13) ^ gf::mul(col[2], 9) ^ gf::mul(col[3], 14);
164    }
165}
166
167/// Encrypt one block in place.
168pub fn encrypt_block(sched: &Schedule, block: &mut [u8]) -> Result<()> {
169    ensure!(block.len() == BLOCK_LEN, InvalidLength, "aes block");
170    let mut s = [0u8; BLOCK_LEN];
171    s.copy_from_slice(block);
172    add_round_key(&mut s, sched.round_key(0));
173    for r in 1..sched.rounds {
174        sub_bytes(&mut s);
175        shift_rows(&mut s);
176        mix_columns(&mut s);
177        add_round_key(&mut s, sched.round_key(r));
178    }
179    sub_bytes(&mut s);
180    shift_rows(&mut s);
181    add_round_key(&mut s, sched.round_key(sched.rounds));
182    block.copy_from_slice(&s);
183    s.zeroize();
184    Ok(())
185}
186
187/// Decrypt one block in place.
188pub fn decrypt_block(sched: &Schedule, block: &mut [u8]) -> Result<()> {
189    ensure!(block.len() == BLOCK_LEN, InvalidLength, "aes block");
190    let mut s = [0u8; BLOCK_LEN];
191    s.copy_from_slice(block);
192    add_round_key(&mut s, sched.round_key(sched.rounds));
193    for r in (1..sched.rounds).rev() {
194        inv_shift_rows(&mut s);
195        inv_sub_bytes(&mut s);
196        add_round_key(&mut s, sched.round_key(r));
197        inv_mix_columns(&mut s);
198    }
199    inv_shift_rows(&mut s);
200    inv_sub_bytes(&mut s);
201    add_round_key(&mut s, sched.round_key(0));
202    block.copy_from_slice(&s);
203    s.zeroize();
204    Ok(())
205}