ic_cipher/aes/
portable.rs1use crate::gf;
8use ic_core::{ensure, Result, Zeroize};
9
10pub const BLOCK_LEN: usize = 16;
12
13pub(crate) const MAX_ROUND_KEYS: usize = 15 * BLOCK_LEN;
15
16fn 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#[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 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 #[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#[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
167pub 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
187pub 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}