1use aes::Aes128;
25use aes::cipher::{BlockCipherEncrypt, KeyInit};
26use sheathe_core::{Error, Result};
27
28mod pssh;
29pub use pssh::ProtectionSystem;
30
31#[derive(Debug, Clone, Copy, PartialEq, Eq)]
33pub enum Scheme {
34 Cenc,
36 Cens,
38 Cbc1,
40 Cbcs,
42}
43
44impl Scheme {
45 pub fn scheme_type(self) -> [u8; 4] {
47 match self {
48 Scheme::Cenc => *b"cenc",
49 Scheme::Cens => *b"cens",
50 Scheme::Cbc1 => *b"cbc1",
51 Scheme::Cbcs => *b"cbcs",
52 }
53 }
54
55 pub fn is_cbc(self) -> bool {
58 matches!(self, Scheme::Cbc1 | Scheme::Cbcs)
59 }
60
61 pub fn is_pattern(self) -> bool {
64 matches!(self, Scheme::Cens | Scheme::Cbcs)
65 }
66
67 pub fn uses_constant_iv(self) -> bool {
70 matches!(self, Scheme::Cbcs)
71 }
72}
73
74#[derive(Debug, Clone, Copy, PartialEq, Eq)]
79pub struct Pattern {
80 pub crypt_blocks: u8,
82 pub skip_blocks: u8,
84}
85
86impl Pattern {
87 pub const NONE: Pattern = Pattern { crypt_blocks: 0, skip_blocks: 0 };
89 pub const VIDEO: Pattern = Pattern { crypt_blocks: 1, skip_blocks: 9 };
91
92 fn is_patterned(self) -> bool {
94 self.crypt_blocks != 0
95 }
96}
97
98#[derive(Debug, Clone)]
100pub struct ContentKey {
101 pub kid: [u8; 16],
103 pub key: [u8; 16],
105}
106
107impl ContentKey {
108 pub fn rotated(&self, period: u32) -> ContentKey {
112 let n = (period % 16) as usize;
113 let mut kid = [0u8; 16];
114 let mut key = [0u8; 16];
115 for i in 0..16 {
116 kid[i] = self.kid[(i + n) % 16];
117 key[i] = self.key[(i + n) % 16];
118 }
119 ContentKey { kid, key }
120 }
121}
122
123#[derive(Debug, Clone, Copy, PartialEq, Eq)]
126pub struct Subsample {
127 pub clear: u32,
129 pub protected: u32,
131}
132
133pub struct Encryptor {
135 cipher: Aes128,
136}
137
138impl Encryptor {
139 pub fn new(key: &[u8; 16]) -> Self {
141 Self { cipher: Aes128::new_from_slice(key).expect("AES-128 key is 16 bytes") }
142 }
143
144 pub fn encrypt(
149 &self,
150 scheme: Scheme,
151 pattern: Pattern,
152 iv: &[u8; 16],
153 data: &mut [u8],
154 subsamples: &[Subsample],
155 ) -> Result<()> {
156 let total: u64 =
158 subsamples.iter().map(|s| u64::from(s.clear) + u64::from(s.protected)).sum();
159 if total != data.len() as u64 {
160 return Err(Error::malformed("subsample layout does not cover sample"));
161 }
162 if pattern.is_patterned() && !scheme.is_pattern() {
163 return Err(Error::malformed("pattern set on a non-pattern scheme"));
164 }
165 if scheme.is_cbc() {
166 self.cbc(pattern, iv, data, subsamples);
167 } else {
168 self.ctr(pattern, iv, data, subsamples);
169 }
170 Ok(())
171 }
172
173 fn ctr(&self, pattern: Pattern, iv: &[u8; 16], data: &mut [u8], subsamples: &[Subsample]) {
179 let mut counter = *iv;
180 let mut keystream = [0u8; 16];
181 let mut ks_pos = 16usize; for_each_crypt_range(pattern, subsamples, |start, len| {
184 for byte in &mut data[start..start + len] {
185 if ks_pos == 16 {
186 keystream = counter;
187 self.encrypt_block(&mut keystream);
188 incr_be(&mut counter);
189 ks_pos = 0;
190 }
191 *byte ^= keystream[ks_pos];
192 ks_pos += 1;
193 }
194 });
195 }
196
197 fn cbc(&self, pattern: Pattern, iv: &[u8; 16], data: &mut [u8], subsamples: &[Subsample]) {
203 let mut off = 0usize;
204 let mut chain = *iv;
208 for s in subsamples {
209 off += s.clear as usize;
210 if pattern.is_patterned() {
211 chain = *iv;
212 }
213 let mut remaining = s.protected as usize;
214 let mut block_index = 0usize;
215 let cycle = pattern.crypt_blocks as usize + pattern.skip_blocks as usize;
216 while remaining >= 16 {
217 let encrypt =
218 !pattern.is_patterned() || block_index % cycle < pattern.crypt_blocks as usize;
219 if encrypt {
220 let mut block = [0u8; 16];
221 block.copy_from_slice(&data[off..off + 16]);
222 for (b, c) in block.iter_mut().zip(chain.iter()) {
223 *b ^= *c;
224 }
225 self.encrypt_block(&mut block);
226 data[off..off + 16].copy_from_slice(&block);
227 chain = block;
228 }
229 off += 16;
230 remaining -= 16;
231 block_index += 1;
232 }
233 off += remaining; }
235 }
236
237 fn encrypt_block(&self, block: &mut [u8; 16]) {
239 let mut ga = (*block).into();
240 self.cipher.encrypt_block(&mut ga);
241 block.copy_from_slice(&ga);
242 }
243}
244
245fn for_each_crypt_range(
252 pattern: Pattern,
253 subsamples: &[Subsample],
254 mut f: impl FnMut(usize, usize),
255) {
256 let mut off = 0usize;
257 for s in subsamples {
258 off += s.clear as usize;
259 let protected = s.protected as usize;
260 if !pattern.is_patterned() {
261 if protected > 0 {
262 f(off, protected);
263 }
264 off += protected;
265 continue;
266 }
267 let crypt = pattern.crypt_blocks as usize * 16;
268 let skip = pattern.skip_blocks as usize * 16;
269 let mut pos = 0usize;
270 while pos < protected {
271 let phase = crypt.min(protected - pos);
272 let whole = phase - phase % 16;
275 if whole > 0 {
276 f(off + pos, whole);
277 }
278 pos += phase;
279 pos += skip.min(protected - pos);
280 }
281 off += protected;
282 }
283}
284
285fn incr_be(counter: &mut [u8; 16]) {
287 for byte in counter.iter_mut().rev() {
288 let (v, carry) = byte.overflowing_add(1);
289 *byte = v;
290 if !carry {
291 break;
292 }
293 }
294}
295
296#[cfg(test)]
297mod tests {
298 use super::*;
299
300 const KEY: [u8; 16] = [
301 0x2b, 0x7e, 0x15, 0x16, 0x28, 0xae, 0xd2, 0xa6, 0xab, 0xf7, 0x15, 0x88, 0x09, 0xcf, 0x4f,
302 0x3c,
303 ];
304
305 fn hex(s: &str) -> Vec<u8> {
306 (0..s.len()).step_by(2).map(|i| u8::from_str_radix(&s[i..i + 2], 16).unwrap()).collect()
307 }
308
309 #[test]
310 fn cenc_matches_nist_ctr_vector() {
311 let iv = hex("f0f1f2f3f4f5f6f7f8f9fafbfcfdfeff");
313 let mut data = hex("6bc1bee22e409f96e93d7e117393172a");
314 let enc = Encryptor::new(&KEY);
315 let subs = [Subsample { clear: 0, protected: 16 }];
316 enc.encrypt(Scheme::Cenc, Pattern::NONE, iv[..].try_into().unwrap(), &mut data, &subs)
317 .unwrap();
318 assert_eq!(data, hex("874d6191b620e3261bef6864990db6ce"));
319 }
320
321 #[test]
322 fn cbc_schemes_match_nist_cbc_vector() {
323 let iv = hex("000102030405060708090a0b0c0d0e0f");
327 let subs = [Subsample { clear: 0, protected: 16 }];
328 for (scheme, pattern) in [(Scheme::Cbc1, Pattern::NONE), (Scheme::Cbcs, Pattern::VIDEO)] {
329 let mut data = hex("6bc1bee22e409f96e93d7e117393172a");
330 let enc = Encryptor::new(&KEY);
331 enc.encrypt(scheme, pattern, iv[..].try_into().unwrap(), &mut data, &subs).unwrap();
332 assert_eq!(data, hex("7649abac8119b246cee98e9b12e9197d"), "{scheme:?}");
333 }
334 }
335
336 #[test]
337 fn cenc_leaves_clear_bytes_untouched() {
338 let iv = [0u8; 16];
339 let mut data = vec![0xAAu8; 32];
340 let enc = Encryptor::new(&KEY);
341 enc.encrypt(
343 Scheme::Cenc,
344 Pattern::NONE,
345 &iv,
346 &mut data,
347 &[Subsample { clear: 8, protected: 24 }],
348 )
349 .unwrap();
350 assert!(data[..8].iter().all(|&b| b == 0xAA), "clear prefix must be untouched");
351 assert!(data[8..].iter().any(|&b| b != 0xAA), "protected region must change");
352 }
353
354 #[test]
355 fn pattern_schemes_skip_blocks() {
356 let iv = [0u8; 16];
359 for scheme in [Scheme::Cbcs, Scheme::Cens] {
360 let mut data = vec![0x11u8; 160];
361 let original = data.clone();
362 let enc = Encryptor::new(&KEY);
363 enc.encrypt(
364 scheme,
365 Pattern::VIDEO,
366 &iv,
367 &mut data,
368 &[Subsample { clear: 0, protected: 160 }],
369 )
370 .unwrap();
371 assert_ne!(data[..16], original[..16], "{scheme:?}: first block encrypted");
372 assert_eq!(data[16..], original[16..], "{scheme:?}: blocks 1..9 skipped");
373 }
374 }
375
376 #[test]
377 fn rejects_mismatched_layout() {
378 let enc = Encryptor::new(&KEY);
379 let mut data = vec![0u8; 10];
380 let err = enc.encrypt(
381 Scheme::Cenc,
382 Pattern::NONE,
383 &[0u8; 16],
384 &mut data,
385 &[Subsample { clear: 0, protected: 9 }],
386 );
387 assert!(err.is_err());
388 }
389
390 #[test]
391 fn content_key_rotation_left_rotates_by_period() {
392 let base = ContentKey { kid: KEY, key: KEY };
393 assert_eq!(base.rotated(0).kid, KEY, "period 0 is unchanged");
394 let mut expect = KEY;
396 expect.rotate_left(1);
397 assert_eq!(base.rotated(1).kid, expect);
398 assert_eq!(base.rotated(1).key, expect);
399 assert_eq!(base.rotated(16).kid, KEY);
401 }
402
403 #[test]
404 fn rejects_pattern_on_non_pattern_scheme() {
405 let enc = Encryptor::new(&KEY);
406 let mut data = vec![0u8; 16];
407 let err = enc.encrypt(
408 Scheme::Cenc,
409 Pattern::VIDEO,
410 &[0u8; 16],
411 &mut data,
412 &[Subsample { clear: 0, protected: 16 }],
413 );
414 assert!(err.is_err());
415 }
416
417 #[test]
421 fn ctr_schemes_round_trip_across_subsamples() {
422 let enc = Encryptor::new(&KEY);
423 let iv = [3u8; 16];
424 let subs = [Subsample { clear: 5, protected: 40 }, Subsample { clear: 10, protected: 65 }];
425 for (scheme, pattern) in [(Scheme::Cenc, Pattern::NONE), (Scheme::Cens, Pattern::VIDEO)] {
426 let original: Vec<u8> = (0..120u8).collect();
427 let mut data = original.clone();
428 enc.encrypt(scheme, pattern, &iv, &mut data, &subs).unwrap();
429 assert_ne!(data, original, "{scheme:?}: ciphertext must differ");
430 assert_eq!(&data[..5], &original[..5], "{scheme:?}: leading clear bytes preserved");
431 enc.encrypt(scheme, pattern, &iv, &mut data, &subs).unwrap();
432 assert_eq!(data, original, "{scheme:?}: CTR round-trip restores plaintext");
433 }
434 }
435
436 #[test]
439 fn cbc1_leaves_trailing_partial_clear() {
440 let enc = Encryptor::new(&KEY);
441 let iv = [7u8; 16];
442 let original: Vec<u8> = (0..40u8).collect();
444 let mut data = original.clone();
445 enc.encrypt(
446 Scheme::Cbc1,
447 Pattern::NONE,
448 &iv,
449 &mut data,
450 &[Subsample { clear: 3, protected: 37 }],
451 )
452 .unwrap();
453 assert_eq!(&data[..3], &original[..3], "leading clear preserved");
454 assert_ne!(&data[3..35], &original[3..35], "full blocks encrypted");
455 assert_eq!(&data[35..], &original[35..], "trailing partial block left clear");
456 }
457}