1use aes::cipher::generic_array::GenericArray;
18use aes::cipher::{BlockEncrypt, KeyInit};
19use aes::Aes128;
20use sheathe_core::{Error, Result};
21
22#[derive(Debug, Clone, Copy, PartialEq, Eq)]
24pub enum Scheme {
25 Cenc,
27 Cbcs,
29}
30
31impl Scheme {
32 pub fn scheme_type(self) -> [u8; 4] {
34 match self {
35 Scheme::Cenc => *b"cenc",
36 Scheme::Cbcs => *b"cbcs",
37 }
38 }
39}
40
41#[derive(Debug, Clone)]
43pub struct ContentKey {
44 pub kid: [u8; 16],
46 pub key: [u8; 16],
48}
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq)]
53pub struct Subsample {
54 pub clear: u32,
56 pub protected: u32,
58}
59
60const CBCS_CRYPT_BLOCKS: usize = 1;
62const CBCS_PATTERN_BLOCKS: usize = 10;
63
64pub struct Encryptor {
66 cipher: Aes128,
67}
68
69impl Encryptor {
70 pub fn new(key: &[u8; 16]) -> Self {
72 Self {
73 cipher: Aes128::new(GenericArray::from_slice(key)),
74 }
75 }
76
77 pub fn encrypt(
80 &self,
81 scheme: Scheme,
82 iv: &[u8; 16],
83 data: &mut [u8],
84 subsamples: &[Subsample],
85 ) -> Result<()> {
86 let total: u64 = subsamples
88 .iter()
89 .map(|s| u64::from(s.clear) + u64::from(s.protected))
90 .sum();
91 if total != data.len() as u64 {
92 return Err(Error::malformed("subsample layout does not cover sample"));
93 }
94 match scheme {
95 Scheme::Cenc => self.cenc(iv, data, subsamples),
96 Scheme::Cbcs => self.cbcs(iv, data, subsamples),
97 }
98 Ok(())
99 }
100
101 fn cenc(&self, iv: &[u8; 16], data: &mut [u8], subsamples: &[Subsample]) {
103 let mut counter = *iv;
104 let mut keystream = [0u8; 16];
105 let mut ks_pos = 16usize; let mut off = 0usize;
107
108 for s in subsamples {
109 off += s.clear as usize;
110 let end = off + s.protected as usize;
111 while off < end {
112 if ks_pos == 16 {
113 keystream = counter;
114 self.encrypt_block(&mut keystream);
115 incr_be(&mut counter);
116 ks_pos = 0;
117 }
118 data[off] ^= keystream[ks_pos];
119 ks_pos += 1;
120 off += 1;
121 }
122 }
123 }
124
125 fn cbcs(&self, iv: &[u8; 16], data: &mut [u8], subsamples: &[Subsample]) {
127 let mut off = 0usize;
128 for s in subsamples {
129 off += s.clear as usize;
130 let mut remaining = s.protected as usize;
131 let mut chain = *iv;
132 let mut block_index = 0usize;
133 while remaining >= 16 {
134 if block_index % CBCS_PATTERN_BLOCKS < CBCS_CRYPT_BLOCKS {
135 let mut block = [0u8; 16];
136 block.copy_from_slice(&data[off..off + 16]);
137 for (b, c) in block.iter_mut().zip(chain.iter()) {
138 *b ^= *c;
139 }
140 self.encrypt_block(&mut block);
141 data[off..off + 16].copy_from_slice(&block);
142 chain = block;
143 }
144 off += 16;
145 remaining -= 16;
146 block_index += 1;
147 }
148 off += remaining; }
150 }
151
152 fn encrypt_block(&self, block: &mut [u8; 16]) {
154 let mut ga = GenericArray::clone_from_slice(block);
155 self.cipher.encrypt_block(&mut ga);
156 block.copy_from_slice(&ga);
157 }
158}
159
160fn incr_be(counter: &mut [u8; 16]) {
162 for byte in counter.iter_mut().rev() {
163 let (v, carry) = byte.overflowing_add(1);
164 *byte = v;
165 if !carry {
166 break;
167 }
168 }
169}
170
171#[cfg(test)]
172mod tests {
173 use super::*;
174
175 const KEY: [u8; 16] = [
176 0x2b, 0x7e, 0x15, 0x16, 0x28, 0xae, 0xd2, 0xa6, 0xab, 0xf7, 0x15, 0x88, 0x09, 0xcf, 0x4f,
177 0x3c,
178 ];
179
180 fn hex(s: &str) -> Vec<u8> {
181 (0..s.len())
182 .step_by(2)
183 .map(|i| u8::from_str_radix(&s[i..i + 2], 16).unwrap())
184 .collect()
185 }
186
187 #[test]
188 fn cenc_matches_nist_ctr_vector() {
189 let iv = hex("f0f1f2f3f4f5f6f7f8f9fafbfcfdfeff");
191 let mut data = hex("6bc1bee22e409f96e93d7e117393172a");
192 let enc = Encryptor::new(&KEY);
193 let subs = [Subsample {
194 clear: 0,
195 protected: 16,
196 }];
197 enc.encrypt(Scheme::Cenc, iv[..].try_into().unwrap(), &mut data, &subs)
198 .unwrap();
199 assert_eq!(data, hex("874d6191b620e3261bef6864990db6ce"));
200 }
201
202 #[test]
203 fn cbcs_first_block_matches_nist_cbc_vector() {
204 let iv = hex("000102030405060708090a0b0c0d0e0f");
206 let mut data = hex("6bc1bee22e409f96e93d7e117393172a");
207 let enc = Encryptor::new(&KEY);
208 let subs = [Subsample {
209 clear: 0,
210 protected: 16,
211 }];
212 enc.encrypt(Scheme::Cbcs, iv[..].try_into().unwrap(), &mut data, &subs)
213 .unwrap();
214 assert_eq!(data, hex("7649abac8119b246cee98e9b12e9197d"));
215 }
216
217 #[test]
218 fn cenc_leaves_clear_bytes_untouched() {
219 let iv = [0u8; 16];
220 let mut data = vec![0xAAu8; 32];
221 let enc = Encryptor::new(&KEY);
222 enc.encrypt(
224 Scheme::Cenc,
225 &iv,
226 &mut data,
227 &[Subsample {
228 clear: 8,
229 protected: 24,
230 }],
231 )
232 .unwrap();
233 assert!(
234 data[..8].iter().all(|&b| b == 0xAA),
235 "clear prefix must be untouched"
236 );
237 assert!(
238 data[8..].iter().any(|&b| b != 0xAA),
239 "protected region must change"
240 );
241 }
242
243 #[test]
244 fn cbcs_pattern_skips_blocks() {
245 let iv = [0u8; 16];
246 let mut data = vec![0x11u8; 160];
248 let original = data.clone();
249 let enc = Encryptor::new(&KEY);
250 enc.encrypt(
251 Scheme::Cbcs,
252 &iv,
253 &mut data,
254 &[Subsample {
255 clear: 0,
256 protected: 160,
257 }],
258 )
259 .unwrap();
260 assert_ne!(data[..16], original[..16], "first block encrypted");
261 assert_eq!(data[16..], original[16..], "blocks 1..9 skipped");
262 }
263
264 #[test]
265 fn rejects_mismatched_layout() {
266 let enc = Encryptor::new(&KEY);
267 let mut data = vec![0u8; 10];
268 let err = enc.encrypt(
269 Scheme::Cenc,
270 &[0u8; 16],
271 &mut data,
272 &[Subsample {
273 clear: 0,
274 protected: 9,
275 }],
276 );
277 assert!(err.is_err());
278 }
279
280 #[test]
281 fn cenc_round_trips_across_subsamples() {
282 let enc = Encryptor::new(&KEY);
286 let iv = [3u8; 16];
287 let subs = [
288 Subsample {
289 clear: 5,
290 protected: 20,
291 },
292 Subsample {
293 clear: 10,
294 protected: 65,
295 },
296 ];
297 let original: Vec<u8> = (0..100u8).collect();
298 let mut data = original.clone();
299
300 enc.encrypt(Scheme::Cenc, &iv, &mut data, &subs).unwrap();
301 assert_ne!(data, original, "ciphertext must differ");
302 assert_eq!(&data[..5], &original[..5], "leading clear bytes preserved");
303
304 enc.encrypt(Scheme::Cenc, &iv, &mut data, &subs).unwrap();
305 assert_eq!(data, original, "CTR round-trip restores plaintext");
306 }
307}