1use crate::aes::BLOCK_LEN;
9use ic_core::traits::BlockCipher;
10use ic_core::{ensure, Result, Zeroize};
11
12pub fn ctr_xor<C: BlockCipher>(cipher: &C, iv: &[u8], data: &mut [u8]) -> Result<()> {
18 ensure!(
19 iv.len() == BLOCK_LEN,
20 InvalidLength,
21 "ctr iv must be 16 bytes"
22 );
23 let mut counter = [0u8; BLOCK_LEN];
24 counter.copy_from_slice(iv);
25
26 let mut keystream = [0u8; BLOCK_LEN * CTR_BATCH];
30
31 for chunk in data.chunks_mut(BLOCK_LEN * CTR_BATCH) {
32 let blocks = chunk.len().div_ceil(BLOCK_LEN);
33 for i in 0..blocks {
34 keystream[i * BLOCK_LEN..(i + 1) * BLOCK_LEN].copy_from_slice(&counter);
35 increment_be(&mut counter);
36 }
37 cipher.encrypt_blocks(&mut keystream[..blocks * BLOCK_LEN])?;
38 for (d, k) in chunk.iter_mut().zip(keystream.iter()) {
39 *d ^= k;
40 }
41 }
42 keystream.zeroize();
43 Ok(())
44}
45
46const CTR_BATCH: usize = 8;
51
52#[inline]
54pub fn increment_be(counter: &mut [u8]) {
55 for byte in counter.iter_mut().rev() {
56 let (v, carry) = byte.overflowing_add(1);
57 *byte = v;
58 if !carry {
59 break;
60 }
61 }
62}
63
64#[inline]
66pub fn increment_be32(counter: &mut [u8; BLOCK_LEN]) {
67 let mut n = u32::from_be_bytes([counter[12], counter[13], counter[14], counter[15]]);
68 n = n.wrapping_add(1);
69 counter[12..].copy_from_slice(&n.to_be_bytes());
70}
71
72pub fn cbc_encrypt<C: BlockCipher>(cipher: &C, iv: &[u8], data: &mut [u8]) -> Result<()> {
76 ensure!(
77 iv.len() == BLOCK_LEN,
78 InvalidLength,
79 "cbc iv must be 16 bytes"
80 );
81 ensure!(
82 data.len() % BLOCK_LEN == 0,
83 InvalidLength,
84 "cbc input must be block-aligned"
85 );
86 let mut prev = [0u8; BLOCK_LEN];
87 prev.copy_from_slice(iv);
88 for block in data.chunks_mut(BLOCK_LEN) {
89 for (b, p) in block.iter_mut().zip(prev.iter()) {
90 *b ^= p;
91 }
92 cipher.encrypt_block(block)?;
93 prev.copy_from_slice(block);
94 }
95 Ok(())
96}
97
98pub fn cbc_decrypt<C: BlockCipher>(cipher: &C, iv: &[u8], data: &mut [u8]) -> Result<()> {
100 ensure!(
101 iv.len() == BLOCK_LEN,
102 InvalidLength,
103 "cbc iv must be 16 bytes"
104 );
105 ensure!(
106 data.len() % BLOCK_LEN == 0,
107 InvalidLength,
108 "cbc input must be block-aligned"
109 );
110 let mut prev = [0u8; BLOCK_LEN];
111 prev.copy_from_slice(iv);
112 let mut saved = [0u8; BLOCK_LEN];
113 for block in data.chunks_mut(BLOCK_LEN) {
114 saved.copy_from_slice(block);
115 cipher.decrypt_block(block)?;
116 for (b, p) in block.iter_mut().zip(prev.iter()) {
117 *b ^= p;
118 }
119 prev.copy_from_slice(&saved);
120 }
121 saved.zeroize();
122 Ok(())
123}
124
125pub fn pkcs7_pad(buf: &mut [u8], len: usize) -> Result<usize> {
129 let pad = BLOCK_LEN - (len % BLOCK_LEN);
130 ensure!(
131 len + pad <= buf.len(),
132 InvalidLength,
133 "pkcs7 padding buffer"
134 );
135 for b in buf[len..len + pad].iter_mut() {
136 *b = pad as u8;
137 }
138 Ok(len + pad)
139}
140
141pub fn pkcs7_unpad(buf: &[u8]) -> Result<usize> {
147 ensure!(
148 !buf.is_empty() && buf.len() % BLOCK_LEN == 0,
149 InvalidLength,
150 "pkcs7 input must be block-aligned"
151 );
152 let pad = buf[buf.len() - 1];
153 let in_range = ((pad.wrapping_sub(1)) < BLOCK_LEN as u8) as u8;
155 let mut bad = in_range ^ 1;
156 for i in 0..BLOCK_LEN {
157 let idx = buf.len() - BLOCK_LEN + i;
158 let is_pad_byte = (((pad as i16) - ((BLOCK_LEN - i) as i16)) >= 0) as u8;
160 bad |= (buf[idx] ^ pad) & is_pad_byte.wrapping_neg();
161 }
162 ensure!(bad == 0, MalformedEncoding, "pkcs7 padding");
163 Ok(buf.len() - pad as usize)
164}
165
166#[cfg(test)]
167mod tests {
168 use super::*;
169 use crate::aes::Aes128;
170 use ic_core::codec::{hex, unhex};
171
172 const SP_KEY: &str = "2b7e151628aed2a6abf7158809cf4f3c";
173 const SP_IV: &str = "000102030405060708090a0b0c0d0e0f";
174 const SP_PT: &str = "6bc1bee22e409f96e93d7e117393172a\
176 ae2d8a571e03ac9c9eb76fac45af8e51\
177 30c81c46a35ce411e5fbc1191a0a52ef\
178 f69f2445df4f9b17ad2b417be66c3710";
179
180 #[test]
181 fn sp800_38a_cbc_vector() {
182 let c = Aes128::new(&unhex(SP_KEY).unwrap()).unwrap();
183 let mut data = unhex(SP_PT).unwrap();
184 cbc_encrypt(&c, &unhex(SP_IV).unwrap(), &mut data).unwrap();
185 assert_eq!(
186 hex(&data),
187 "7649abac8119b246cee98e9b12e9197d\
188 5086cb9b507219ee95db113a917678b2\
189 73bed6b8e3c1743b7116e69e22229516\
190 3ff1caa1681fac09120eca307586e1a7"
191 .replace(char::is_whitespace, "")
192 );
193 cbc_decrypt(&c, &unhex(SP_IV).unwrap(), &mut data).unwrap();
194 assert_eq!(hex(&data), SP_PT.replace(char::is_whitespace, ""));
195 }
196
197 #[test]
198 fn sp800_38a_ctr_vector() {
199 let c = Aes128::new(&unhex(SP_KEY).unwrap()).unwrap();
200 let iv = unhex("f0f1f2f3f4f5f6f7f8f9fafbfcfdfeff").unwrap();
201 let mut data = unhex(SP_PT).unwrap();
202 ctr_xor(&c, &iv, &mut data).unwrap();
203 assert_eq!(
204 hex(&data),
205 "874d6191b620e3261bef6864990db6ce\
206 9806f66b7970fdff8617187bb9fffdff\
207 5ae4df3edbd5d35e5b4f09020db03eab\
208 1e031dda2fbe03d1792170a0f3009cee"
209 .replace(char::is_whitespace, "")
210 );
211 ctr_xor(&c, &iv, &mut data).unwrap();
213 assert_eq!(hex(&data), SP_PT.replace(char::is_whitespace, ""));
214 }
215
216 #[test]
217 fn ctr_handles_partial_final_block() {
218 let c = Aes128::new(&[0u8; 16]).unwrap();
219 let mut data = [0u8; 37];
220 ctr_xor(&c, &[0u8; 16], &mut data).unwrap();
221 let encrypted = data;
222 ctr_xor(&c, &[0u8; 16], &mut data).unwrap();
223 assert_eq!(data, [0u8; 37]);
224 assert_ne!(encrypted, [0u8; 37]);
225 }
226
227 #[test]
228 fn counter_increment_carries() {
229 let mut c = [0xffu8; 16];
230 increment_be(&mut c);
231 assert_eq!(c, [0u8; 16]);
232 let mut c = [0u8; 16];
233 c[15] = 0xff;
234 increment_be(&mut c);
235 assert_eq!(c[14], 1);
236 assert_eq!(c[15], 0);
237 }
238
239 #[test]
240 fn gcm_counter_wraps_only_low_32_bits() {
241 let mut c = [0u8; 16];
242 c[11] = 0x7f;
243 c[12..].copy_from_slice(&0xffff_ffffu32.to_be_bytes());
244 increment_be32(&mut c);
245 assert_eq!(&c[12..], &[0, 0, 0, 0]);
246 assert_eq!(
247 c[11], 0x7f,
248 "carry must not propagate past the counter field"
249 );
250 }
251
252 #[test]
253 fn pkcs7_roundtrip_including_full_block() {
254 for len in 0..33usize {
255 let mut buf = vec![0xAAu8; len + BLOCK_LEN];
256 let padded = pkcs7_pad(&mut buf, len).unwrap();
257 assert_eq!(padded % BLOCK_LEN, 0);
258 assert_eq!(pkcs7_unpad(&buf[..padded]).unwrap(), len, "len {len}");
259 }
260 }
261
262 #[test]
263 fn pkcs7_rejects_corrupt_padding() {
264 let mut buf = [0u8; 16];
265 let n = pkcs7_pad(&mut buf, 8).unwrap();
266 buf[n - 2] ^= 1;
267 assert!(pkcs7_unpad(&buf[..n]).is_err());
268 let mut zero = [0u8; 16];
269 zero[15] = 0;
270 assert!(pkcs7_unpad(&zero).is_err());
271 let mut big = [0u8; 16];
272 big[15] = 17;
273 assert!(pkcs7_unpad(&big).is_err());
274 }
275}