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