use aes::cipher::{BlockCipherEncrypt, KeyInit};
use aes::{Aes128, Aes256};
use sheathe_core::{Error, Result};
mod pssh;
pub use pssh::ProtectionSystem;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Scheme {
Cenc,
Cens,
Cbc1,
Cbcs,
}
impl Scheme {
pub fn scheme_type(self) -> [u8; 4] {
match self {
Scheme::Cenc => *b"cenc",
Scheme::Cens => *b"cens",
Scheme::Cbc1 => *b"cbc1",
Scheme::Cbcs => *b"cbcs",
}
}
pub fn is_cbc(self) -> bool {
matches!(self, Scheme::Cbc1 | Scheme::Cbcs)
}
pub fn is_pattern(self) -> bool {
matches!(self, Scheme::Cens | Scheme::Cbcs)
}
pub fn uses_constant_iv(self) -> bool {
matches!(self, Scheme::Cbcs)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Pattern {
pub crypt_blocks: u8,
pub skip_blocks: u8,
}
impl Pattern {
pub const NONE: Pattern = Pattern { crypt_blocks: 0, skip_blocks: 0 };
pub const VIDEO: Pattern = Pattern { crypt_blocks: 1, skip_blocks: 9 };
pub const fn from_blocks(crypt: u8, skip: u8) -> Pattern {
Pattern { crypt_blocks: crypt, skip_blocks: skip }
}
fn is_patterned(self) -> bool {
self.crypt_blocks != 0
}
}
#[derive(Debug, Clone)]
pub struct ContentKey {
pub kid: [u8; 16],
pub key: [u8; 16],
}
impl ContentKey {
pub fn rotated(&self, period: u32) -> ContentKey {
let n = (period % 16) as usize;
let mut kid = [0u8; 16];
let mut key = [0u8; 16];
for i in 0..16 {
kid[i] = self.kid[(i + n) % 16];
key[i] = self.key[(i + n) % 16];
}
ContentKey { kid, key }
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Subsample {
pub clear: u32,
pub protected: u32,
}
pub struct Encryptor {
cipher: Aes128,
}
impl Encryptor {
pub fn new(key: &[u8; 16]) -> Self {
Self { cipher: Aes128::new_from_slice(key).expect("AES-128 key is 16 bytes") }
}
pub fn encrypt(
&self,
scheme: Scheme,
pattern: Pattern,
iv: &[u8; 16],
data: &mut [u8],
subsamples: &[Subsample],
) -> Result<()> {
let total: u64 =
subsamples.iter().map(|s| u64::from(s.clear) + u64::from(s.protected)).sum();
if total != data.len() as u64 {
return Err(Error::malformed("subsample layout does not cover sample"));
}
if pattern.is_patterned() && !scheme.is_pattern() {
return Err(Error::malformed("pattern set on a non-pattern scheme"));
}
if scheme.is_cbc() {
self.cbc(pattern, iv, data, subsamples);
} else {
self.ctr(pattern, iv, data, subsamples);
}
Ok(())
}
fn ctr(&self, pattern: Pattern, iv: &[u8; 16], data: &mut [u8], subsamples: &[Subsample]) {
let mut counter = *iv;
let mut keystream = [0u8; 16];
let mut ks_pos = 16usize;
for_each_crypt_range(pattern, subsamples, |start, len| {
for byte in &mut data[start..start + len] {
if ks_pos == 16 {
keystream = counter;
self.encrypt_block(&mut keystream);
incr_be(&mut counter);
ks_pos = 0;
}
*byte ^= keystream[ks_pos];
ks_pos += 1;
}
});
}
fn cbc(&self, pattern: Pattern, iv: &[u8; 16], data: &mut [u8], subsamples: &[Subsample]) {
let mut off = 0usize;
let mut chain = *iv;
for s in subsamples {
off += s.clear as usize;
if pattern.is_patterned() {
chain = *iv;
}
let mut remaining = s.protected as usize;
let mut block_index = 0usize;
let cycle = pattern.crypt_blocks as usize + pattern.skip_blocks as usize;
while remaining >= 16 {
let encrypt =
!pattern.is_patterned() || block_index % cycle < pattern.crypt_blocks as usize;
if encrypt {
let mut block = [0u8; 16];
block.copy_from_slice(&data[off..off + 16]);
for (b, c) in block.iter_mut().zip(chain.iter()) {
*b ^= *c;
}
self.encrypt_block(&mut block);
data[off..off + 16].copy_from_slice(&block);
chain = block;
}
off += 16;
remaining -= 16;
block_index += 1;
}
off += remaining; }
}
fn encrypt_block(&self, block: &mut [u8; 16]) {
let mut ga = (*block).into();
self.cipher.encrypt_block(&mut ga);
block.copy_from_slice(&ga);
}
}
fn for_each_crypt_range(
pattern: Pattern,
subsamples: &[Subsample],
mut f: impl FnMut(usize, usize),
) {
let mut off = 0usize;
for s in subsamples {
off += s.clear as usize;
let protected = s.protected as usize;
if !pattern.is_patterned() {
if protected > 0 {
f(off, protected);
}
off += protected;
continue;
}
let crypt = pattern.crypt_blocks as usize * 16;
let skip = pattern.skip_blocks as usize * 16;
let mut pos = 0usize;
while pos < protected {
let phase = crypt.min(protected - pos);
let whole = phase - phase % 16;
if whole > 0 {
f(off + pos, whole);
}
pos += phase;
pos += skip.min(protected - pos);
}
off += protected;
}
}
fn incr_be(counter: &mut [u8; 16]) {
for byte in counter.iter_mut().rev() {
let (v, carry) = byte.overflowing_add(1);
*byte = v;
if !carry {
break;
}
}
}
pub fn aes_cbc_pkcs7_encrypt(key: &[u8], iv: &[u8], data: &[u8]) -> Result<Vec<u8>> {
if iv.len() != 16 {
return Err(Error::malformed("CBC IV must be 16 bytes"));
}
let padded = pkcs7_pad(data, 16);
aes_cbc_raw(key, iv, &padded)
}
pub fn aes_cbc_pkcs7_decrypt(key: &[u8], iv: &[u8], data: &[u8]) -> Result<Vec<u8>> {
if iv.len() != 16 {
return Err(Error::malformed("CBC IV must be 16 bytes"));
}
if data.len() % 16 != 0 || data.is_empty() {
return Err(Error::malformed("CBC ciphertext must be a non-empty multiple of 16"));
}
let plain = aes_cbc_raw_decrypt(key, iv, data)?;
pkcs7_unpad(&plain)
}
fn pkcs7_pad(data: &[u8], block: usize) -> Vec<u8> {
let n = block - (data.len() % block);
let mut v = data.to_vec();
v.extend(std::iter::repeat_n(n as u8, n));
v
}
fn pkcs7_unpad(data: &[u8]) -> Result<Vec<u8>> {
let n = *data.last().ok_or_else(|| Error::malformed("empty PKCS7 buffer"))? as usize;
if n == 0
|| n > 16
|| n > data.len()
|| !data[data.len() - n..].iter().all(|&b| b as usize == n)
{
return Err(Error::malformed("invalid PKCS7 padding"));
}
Ok(data[..data.len() - n].to_vec())
}
fn aes_cbc_raw(key: &[u8], iv: &[u8], data: &[u8]) -> Result<Vec<u8>> {
let mut prev = [0u8; 16];
prev.copy_from_slice(iv);
let mut out = data.to_vec();
for chunk in out.chunks_mut(16) {
for (b, p) in chunk.iter_mut().zip(prev.iter()) {
*b ^= *p;
}
encrypt_block(key, chunk)?;
prev.copy_from_slice(chunk);
}
Ok(out)
}
fn aes_cbc_raw_decrypt(key: &[u8], iv: &[u8], data: &[u8]) -> Result<Vec<u8>> {
let mut prev = [0u8; 16];
prev.copy_from_slice(iv);
let mut out = data.to_vec();
for chunk in out.chunks_mut(16) {
let saved = <[u8; 16]>::try_from(&chunk[..]).unwrap();
decrypt_block(key, chunk)?;
for (b, p) in chunk.iter_mut().zip(prev.iter()) {
*b ^= *p;
}
prev = saved;
}
Ok(out)
}
fn encrypt_block(key: &[u8], block: &mut [u8]) -> Result<()> {
match key.len() {
16 => {
let cipher = Aes128::new_from_slice(key).expect("16");
let mut b = [0u8; 16];
b.copy_from_slice(block);
cipher.encrypt_block((&mut b).into());
block.copy_from_slice(&b);
Ok(())
}
32 => {
let cipher = Aes256::new_from_slice(key).expect("32");
let mut b = [0u8; 16];
b.copy_from_slice(block);
cipher.encrypt_block((&mut b).into());
block.copy_from_slice(&b);
Ok(())
}
_ => Err(Error::malformed("AES key must be 16 or 32 bytes")),
}
}
fn decrypt_block(key: &[u8], block: &mut [u8]) -> Result<()> {
use aes::cipher::BlockCipherDecrypt;
match key.len() {
16 => {
let cipher = Aes128::new_from_slice(key).expect("16");
let mut b = [0u8; 16];
b.copy_from_slice(block);
cipher.decrypt_block((&mut b).into());
block.copy_from_slice(&b);
Ok(())
}
32 => {
let cipher = Aes256::new_from_slice(key).expect("32");
let mut b = [0u8; 16];
b.copy_from_slice(block);
cipher.decrypt_block((&mut b).into());
block.copy_from_slice(&b);
Ok(())
}
_ => Err(Error::malformed("AES key must be 16 or 32 bytes")),
}
}
#[cfg(test)]
mod tests {
use super::*;
const KEY: [u8; 16] = [
0x2b, 0x7e, 0x15, 0x16, 0x28, 0xae, 0xd2, 0xa6, 0xab, 0xf7, 0x15, 0x88, 0x09, 0xcf, 0x4f,
0x3c,
];
fn hex(s: &str) -> Vec<u8> {
(0..s.len()).step_by(2).map(|i| u8::from_str_radix(&s[i..i + 2], 16).unwrap()).collect()
}
#[test]
fn cenc_matches_nist_ctr_vector() {
let iv = hex("f0f1f2f3f4f5f6f7f8f9fafbfcfdfeff");
let mut data = hex("6bc1bee22e409f96e93d7e117393172a");
let enc = Encryptor::new(&KEY);
let subs = [Subsample { clear: 0, protected: 16 }];
enc.encrypt(Scheme::Cenc, Pattern::NONE, iv[..].try_into().unwrap(), &mut data, &subs)
.unwrap();
assert_eq!(data, hex("874d6191b620e3261bef6864990db6ce"));
}
#[test]
fn cbc_schemes_match_nist_cbc_vector() {
let iv = hex("000102030405060708090a0b0c0d0e0f");
let subs = [Subsample { clear: 0, protected: 16 }];
for (scheme, pattern) in [(Scheme::Cbc1, Pattern::NONE), (Scheme::Cbcs, Pattern::VIDEO)] {
let mut data = hex("6bc1bee22e409f96e93d7e117393172a");
let enc = Encryptor::new(&KEY);
enc.encrypt(scheme, pattern, iv[..].try_into().unwrap(), &mut data, &subs).unwrap();
assert_eq!(data, hex("7649abac8119b246cee98e9b12e9197d"), "{scheme:?}");
}
}
#[test]
fn cenc_leaves_clear_bytes_untouched() {
let iv = [0u8; 16];
let mut data = vec![0xAAu8; 32];
let enc = Encryptor::new(&KEY);
enc.encrypt(
Scheme::Cenc,
Pattern::NONE,
&iv,
&mut data,
&[Subsample { clear: 8, protected: 24 }],
)
.unwrap();
assert!(data[..8].iter().all(|&b| b == 0xAA), "clear prefix must be untouched");
assert!(data[8..].iter().any(|&b| b != 0xAA), "protected region must change");
}
#[test]
fn pattern_schemes_skip_blocks() {
let iv = [0u8; 16];
for scheme in [Scheme::Cbcs, Scheme::Cens] {
let mut data = vec![0x11u8; 160];
let original = data.clone();
let enc = Encryptor::new(&KEY);
enc.encrypt(
scheme,
Pattern::VIDEO,
&iv,
&mut data,
&[Subsample { clear: 0, protected: 160 }],
)
.unwrap();
assert_ne!(data[..16], original[..16], "{scheme:?}: first block encrypted");
assert_eq!(data[16..], original[16..], "{scheme:?}: blocks 1..9 skipped");
}
}
#[test]
fn rejects_mismatched_layout() {
let enc = Encryptor::new(&KEY);
let mut data = vec![0u8; 10];
let err = enc.encrypt(
Scheme::Cenc,
Pattern::NONE,
&[0u8; 16],
&mut data,
&[Subsample { clear: 0, protected: 9 }],
);
assert!(err.is_err());
}
#[test]
fn content_key_rotation_left_rotates_by_period() {
let base = ContentKey { kid: KEY, key: KEY };
assert_eq!(base.rotated(0).kid, KEY, "period 0 is unchanged");
let mut expect = KEY;
expect.rotate_left(1);
assert_eq!(base.rotated(1).kid, expect);
assert_eq!(base.rotated(1).key, expect);
assert_eq!(base.rotated(16).kid, KEY);
}
#[test]
fn rejects_pattern_on_non_pattern_scheme() {
let enc = Encryptor::new(&KEY);
let mut data = vec![0u8; 16];
let err = enc.encrypt(
Scheme::Cenc,
Pattern::VIDEO,
&[0u8; 16],
&mut data,
&[Subsample { clear: 0, protected: 16 }],
);
assert!(err.is_err());
}
#[test]
fn ctr_schemes_round_trip_across_subsamples() {
let enc = Encryptor::new(&KEY);
let iv = [3u8; 16];
let subs = [Subsample { clear: 5, protected: 40 }, Subsample { clear: 10, protected: 65 }];
for (scheme, pattern) in [(Scheme::Cenc, Pattern::NONE), (Scheme::Cens, Pattern::VIDEO)] {
let original: Vec<u8> = (0..120u8).collect();
let mut data = original.clone();
enc.encrypt(scheme, pattern, &iv, &mut data, &subs).unwrap();
assert_ne!(data, original, "{scheme:?}: ciphertext must differ");
assert_eq!(&data[..5], &original[..5], "{scheme:?}: leading clear bytes preserved");
enc.encrypt(scheme, pattern, &iv, &mut data, &subs).unwrap();
assert_eq!(data, original, "{scheme:?}: CTR round-trip restores plaintext");
}
}
#[test]
fn cbc1_leaves_trailing_partial_clear() {
let enc = Encryptor::new(&KEY);
let iv = [7u8; 16];
let original: Vec<u8> = (0..40u8).collect();
let mut data = original.clone();
enc.encrypt(
Scheme::Cbc1,
Pattern::NONE,
&iv,
&mut data,
&[Subsample { clear: 3, protected: 37 }],
)
.unwrap();
assert_eq!(&data[..3], &original[..3], "leading clear preserved");
assert_ne!(&data[3..35], &original[3..35], "full blocks encrypted");
assert_eq!(&data[35..], &original[35..], "trailing partial block left clear");
}
}