use aead::array::Array;
use aead::consts::{U19, U24, U32};
use aead::{Aead, Generate, KeyInit, Payload};
use aead_stream::{DecryptorBE32, EncryptorBE32, NewStream, StreamBE32, StreamPrimitive};
use chacha20poly1305::XChaCha20Poly1305;
use gwk_domain::blob::BLOB_CHUNK_BYTES;
use gwk_domain::port::BlobError;
use sha2::{Digest, Sha256};
use zeroize::Zeroize;
pub const MAGIC: [u8; 8] = *b"GWKBLOB\0";
pub const FORMAT_VERSION: u16 = 1;
pub const DEK_BYTES: usize = 32;
pub const WRAP_NONCE_BYTES: usize = 24;
pub const TAG_BYTES: usize = 16;
pub const WRAPPED_DEK_BYTES: usize = DEK_BYTES + TAG_BYTES;
pub const STREAM_NONCE_BYTES: usize = WRAP_NONCE_BYTES - 5;
pub const CHUNK_LEN_BYTES: usize = 4;
pub const MAX_CIPHERTEXT_CHUNK_BYTES: usize = BLOB_CHUNK_BYTES + TAG_BYTES;
pub const FRAMED_CHUNK_BYTES: u64 = (CHUNK_LEN_BYTES + MAX_CIPHERTEXT_CHUNK_BYTES) as u64;
const DIGEST_BYTES: usize = 32;
const FIXED_HEADER_BYTES: usize = MAGIC.len() + 2 + DIGEST_BYTES + 8 + STREAM_NONCE_BYTES + 2 + 2;
pub fn header_len(media_type: &str, kek_id: &str) -> usize {
FIXED_HEADER_BYTES + media_type.len() + kek_id.len()
}
pub fn chunk_count(byte_size: u64) -> u64 {
byte_size.div_ceil(BLOB_CHUNK_BYTES as u64).max(1)
}
pub fn chunk_offset(header_len: usize, index: u64) -> u64 {
header_len as u64 + index * FRAMED_CHUNK_BYTES
}
type Stream = StreamBE32<XChaCha20Poly1305>;
type StreamNonce = aead_stream::Nonce<XChaCha20Poly1305, Stream>;
fn integrity(reason: impl Into<String>) -> BlobError {
BlobError::Integrity(reason.into())
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Header {
pub digest: [u8; DIGEST_BYTES],
pub byte_size: u64,
pub media_type: String,
pub kek_id: String,
pub stream_nonce: [u8; STREAM_NONCE_BYTES],
}
impl Header {
pub fn encode(&self) -> Vec<u8> {
let media = self.media_type.as_bytes();
let kek = self.kek_id.as_bytes();
let mut out = Vec::with_capacity(FIXED_HEADER_BYTES + media.len() + kek.len());
out.extend_from_slice(&MAGIC);
out.extend_from_slice(&FORMAT_VERSION.to_be_bytes());
out.extend_from_slice(&self.digest);
out.extend_from_slice(&self.byte_size.to_be_bytes());
out.extend_from_slice(&self.stream_nonce);
out.extend_from_slice(&(media.len() as u16).to_be_bytes());
out.extend_from_slice(&(kek.len() as u16).to_be_bytes());
out.extend_from_slice(media);
out.extend_from_slice(kek);
out
}
pub fn decode(bytes: &[u8]) -> Result<(Self, usize), BlobError> {
if bytes.len() < FIXED_HEADER_BYTES {
return Err(integrity("container is shorter than a header"));
}
const VERSION_AT: usize = MAGIC.len();
const DIGEST_AT: usize = VERSION_AT + 2;
const SIZE_AT: usize = DIGEST_AT + DIGEST_BYTES;
const NONCE_AT: usize = SIZE_AT + 8;
const MEDIA_LEN_AT: usize = NONCE_AT + STREAM_NONCE_BYTES;
const KEK_LEN_AT: usize = MEDIA_LEN_AT + 2;
if bytes[..MAGIC.len()] != MAGIC {
return Err(integrity("not a gwk blob container"));
}
let version = u16::from_be_bytes(
bytes[VERSION_AT..DIGEST_AT]
.try_into()
.map_err(|_| integrity("version"))?,
);
if version != FORMAT_VERSION {
return Err(integrity(format!(
"container format version {version}, expected {FORMAT_VERSION}"
)));
}
let digest: [u8; DIGEST_BYTES] = bytes[DIGEST_AT..SIZE_AT]
.try_into()
.map_err(|_| integrity("digest"))?;
let byte_size = u64::from_be_bytes(
bytes[SIZE_AT..NONCE_AT]
.try_into()
.map_err(|_| integrity("size"))?,
);
let stream_nonce: [u8; STREAM_NONCE_BYTES] = bytes[NONCE_AT..MEDIA_LEN_AT]
.try_into()
.map_err(|_| integrity("nonce"))?;
let media_len = usize::from(u16::from_be_bytes(
bytes[MEDIA_LEN_AT..KEK_LEN_AT]
.try_into()
.map_err(|_| integrity("len"))?,
));
let kek_len = usize::from(u16::from_be_bytes(
bytes[KEK_LEN_AT..FIXED_HEADER_BYTES]
.try_into()
.map_err(|_| integrity("len"))?,
));
let mut at = FIXED_HEADER_BYTES;
if bytes.len() < at + media_len + kek_len {
return Err(integrity("header runs past the container"));
}
let media_type = std::str::from_utf8(&bytes[at..at + media_len])
.map_err(|_| integrity("media type is not utf-8"))?
.to_owned();
at += media_len;
let kek_id = std::str::from_utf8(&bytes[at..at + kek_len])
.map_err(|_| integrity("kek id is not utf-8"))?
.to_owned();
at += kek_len;
Ok((
Self {
digest,
byte_size,
media_type,
kek_id,
stream_nonce,
},
at,
))
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Sealed {
pub container: Vec<u8>,
pub wrap_nonce: [u8; WRAP_NONCE_BYTES],
pub wrapped_dek: [u8; WRAPPED_DEK_BYTES],
pub digest_hex: String,
}
pub fn seal_with(
plaintext: &[u8],
media_type: &str,
kek: &[u8; DEK_BYTES],
kek_id: &str,
dek: &[u8; DEK_BYTES],
wrap_nonce: &[u8; WRAP_NONCE_BYTES],
stream_nonce: &[u8; STREAM_NONCE_BYTES],
) -> Result<Sealed, BlobError> {
if media_type.len() > u16::MAX as usize || kek_id.len() > u16::MAX as usize {
return Err(integrity("media type or kek id is too long to encode"));
}
let digest: [u8; DIGEST_BYTES] = Sha256::digest(plaintext).into();
let header = Header {
digest,
byte_size: plaintext.len() as u64,
media_type: media_type.to_owned(),
kek_id: kek_id.to_owned(),
stream_nonce: *stream_nonce,
};
let aad = header.encode();
let wrapped_dek = wrap_dek(dek, kek, wrap_nonce, &aad)?;
let mut container = aad.clone();
let mut encryptor =
EncryptorBE32::<XChaCha20Poly1305>::new(&Array(*dek), &StreamNonce::from(*stream_nonce));
let mut chunks: Vec<&[u8]> = plaintext.chunks(BLOB_CHUNK_BYTES).collect();
if chunks.is_empty() {
chunks.push(&[]);
}
let Some((last, rest)) = chunks.split_last() else {
return Err(integrity("no chunk to seal"));
};
for chunk in rest {
let sealed = encryptor
.encrypt_next(Payload {
msg: chunk,
aad: &aad,
})
.map_err(|_| integrity("chunk encryption failed"))?;
push_chunk(&mut container, &sealed)?;
}
let sealed = encryptor
.encrypt_last(Payload {
msg: last,
aad: &aad,
})
.map_err(|_| integrity("final chunk encryption failed"))?;
push_chunk(&mut container, &sealed)?;
Ok(Sealed {
container,
wrap_nonce: *wrap_nonce,
wrapped_dek,
digest_hex: hex_lower(&digest),
})
}
pub fn seal(
plaintext: &[u8],
media_type: &str,
kek: &[u8; DEK_BYTES],
kek_id: &str,
) -> Result<Sealed, BlobError> {
let mut dek: Array<u8, U32> = generate()?;
let wrap_nonce: Array<u8, U24> = generate()?;
let stream_nonce: Array<u8, U19> = generate()?;
let sealed = seal_with(
plaintext,
media_type,
kek,
kek_id,
&dek.0,
&wrap_nonce.0,
&stream_nonce.0,
);
dek.zeroize();
sealed
}
pub fn open(
container: &[u8],
kek: &[u8; DEK_BYTES],
wrap_nonce: &[u8; WRAP_NONCE_BYTES],
wrapped_dek: &[u8; WRAPPED_DEK_BYTES],
) -> Result<(Header, Vec<u8>), BlobError> {
let (header, header_len) = Header::decode(container)?;
let aad = &container[..header_len];
let mut dek = unwrap_dek(wrapped_dek, kek, wrap_nonce, aad)?;
let mut decryptor = DecryptorBE32::<XChaCha20Poly1305>::new(
&Array(dek),
&StreamNonce::from(header.stream_nonce),
);
dek.zeroize();
let mut framed: Vec<&[u8]> = Vec::new();
let mut at = header_len;
while at < container.len() {
let (chunk, next) = read_chunk(container, at)?;
framed.push(chunk);
at = next;
}
let Some((last, rest)) = framed.split_last() else {
return Err(integrity("container has no chunks"));
};
let mut plaintext = Vec::with_capacity(header.byte_size as usize);
for chunk in rest {
let part = decryptor
.decrypt_next(Payload { msg: chunk, aad })
.map_err(|_| integrity("chunk failed authentication"))?;
plaintext.extend_from_slice(&part);
}
let part = decryptor
.decrypt_last(Payload { msg: last, aad })
.map_err(|_| integrity("final chunk failed authentication: tampered or truncated"))?;
plaintext.extend_from_slice(&part);
if plaintext.len() as u64 != header.byte_size {
return Err(integrity(format!(
"container holds {} plaintext bytes, header declares {}",
plaintext.len(),
header.byte_size
)));
}
let actual: [u8; DIGEST_BYTES] = Sha256::digest(&plaintext).into();
if actual != header.digest {
return Err(integrity("plaintext does not hash to the declared digest"));
}
Ok((header, plaintext))
}
pub fn rewrap(
container: &[u8],
old_kek: &[u8; DEK_BYTES],
new_kek: &[u8; DEK_BYTES],
wrap_nonce: &[u8; WRAP_NONCE_BYTES],
wrapped_dek: &[u8; WRAPPED_DEK_BYTES],
new_wrap_nonce: &[u8; WRAP_NONCE_BYTES],
) -> Result<[u8; WRAPPED_DEK_BYTES], BlobError> {
let (_, header_len) = Header::decode(container)?;
let aad = &container[..header_len];
let mut dek = unwrap_dek(wrapped_dek, old_kek, wrap_nonce, aad)?;
let rewrapped = wrap_dek(&dek, new_kek, new_wrap_nonce, aad);
dek.zeroize();
rewrapped
}
pub fn open_chunk(
aad: &[u8],
dek: &[u8; DEK_BYTES],
stream_nonce: &[u8; STREAM_NONCE_BYTES],
position: u32,
last: bool,
ciphertext: &[u8],
) -> Result<Vec<u8>, BlobError> {
Stream::new(&Array(*dek), &StreamNonce::from(*stream_nonce))
.decrypt(
position,
last,
Payload {
msg: ciphertext,
aad,
},
)
.map_err(|_| integrity("chunk failed authentication: tampered or truncated"))
}
pub(crate) fn wrap_dek(
dek: &[u8; DEK_BYTES],
kek: &[u8; DEK_BYTES],
nonce: &[u8; WRAP_NONCE_BYTES],
aad: &[u8],
) -> Result<[u8; WRAPPED_DEK_BYTES], BlobError> {
let sealed = XChaCha20Poly1305::new(&Array(*kek))
.encrypt(&Array(*nonce), Payload { msg: dek, aad })
.map_err(|_| integrity("wrapping the data key failed"))?;
sealed
.try_into()
.map_err(|_| integrity("wrapped key is the wrong length"))
}
pub(crate) fn unwrap_dek(
wrapped: &[u8; WRAPPED_DEK_BYTES],
kek: &[u8; DEK_BYTES],
nonce: &[u8; WRAP_NONCE_BYTES],
aad: &[u8],
) -> Result<[u8; DEK_BYTES], BlobError> {
let dek = XChaCha20Poly1305::new(&Array(*kek))
.decrypt(
&Array(*nonce),
Payload {
msg: wrapped.as_slice(),
aad,
},
)
.map_err(|_| integrity("unwrapping the data key failed: wrong kek or altered header"))?;
dek.try_into()
.map_err(|_| integrity("unwrapped key is the wrong length"))
}
fn push_chunk(out: &mut Vec<u8>, chunk: &[u8]) -> Result<(), BlobError> {
let len = u32::try_from(chunk.len()).map_err(|_| integrity("chunk is too large to frame"))?;
out.extend_from_slice(&len.to_be_bytes());
out.extend_from_slice(chunk);
Ok(())
}
fn read_chunk(container: &[u8], at: usize) -> Result<(&[u8], usize), BlobError> {
let end = at + 4;
if container.len() < end {
return Err(integrity("truncated chunk length"));
}
let len = u32::from_be_bytes(
container[at..end]
.try_into()
.map_err(|_| integrity("chunk length"))?,
) as usize;
let stop = end
.checked_add(len)
.ok_or_else(|| integrity("chunk length overflows the container"))?;
if container.len() < stop {
return Err(integrity("chunk runs past the container"));
}
Ok((&container[end..stop], stop))
}
pub(crate) fn generate<N: aead::array::ArraySize>() -> Result<Array<u8, N>, BlobError> {
Array::try_generate()
.map_err(|e| BlobError::Storage(format!("system randomness unavailable: {e}")))
}
pub(crate) fn hex_lower(bytes: &[u8]) -> String {
let mut out = String::with_capacity(bytes.len() * 2);
for byte in bytes {
out.push_str(&format!("{byte:02x}"));
}
out
}
#[cfg(test)]
mod tests {
use super::*;
const KEK: [u8; DEK_BYTES] = [0x11; DEK_BYTES];
const DEK: [u8; DEK_BYTES] = [0x22; DEK_BYTES];
const WRAP_NONCE: [u8; WRAP_NONCE_BYTES] = [0x33; WRAP_NONCE_BYTES];
const STREAM_NONCE: [u8; STREAM_NONCE_BYTES] = [0x44; STREAM_NONCE_BYTES];
fn seal_fixture(plaintext: &[u8]) -> Sealed {
seal_with(
plaintext,
"application/json",
&KEK,
"kek-test",
&DEK,
&WRAP_NONCE,
&STREAM_NONCE,
)
.expect("seal")
}
#[test]
fn a_sealed_blob_opens_to_exactly_what_went_in() {
for plaintext in [
b"".to_vec(),
b"one chunk".to_vec(),
vec![0xab; BLOB_CHUNK_BYTES], vec![0xcd; BLOB_CHUNK_BYTES + 1], vec![0xef; BLOB_CHUNK_BYTES * 2 + 7], ] {
let sealed = seal_fixture(&plaintext);
let (header, opened) = open(
&sealed.container,
&KEK,
&sealed.wrap_nonce,
&sealed.wrapped_dek,
)
.expect("open");
assert_eq!(opened, plaintext);
assert_eq!(header.byte_size, plaintext.len() as u64);
assert_eq!(header.media_type, "application/json");
assert_eq!(header.kek_id, "kek-test");
}
}
#[test]
fn the_format_is_pinned_byte_for_byte() {
let sealed = seal_fixture(b"gridwork");
let hex: String = sealed
.container
.iter()
.map(|b| format!("{b:02x}"))
.collect();
assert_eq!(
hex,
"47574b424c4f4200000143e27ffff32e033107110d7303fa2fcf4eb11ee1cfe54f5017f9a653\
084c5d1d00000000000000084444444444444444444444444444444444444400100008617070\
6c69636174696f6e2f6a736f6e6b656b2d7465737400000018658bce0c2761f96df661755b52\
019f8186b90db3313e5722"
.replace([' ', '\n'], "")
);
assert_eq!(
sealed.digest_hex,
"43e27ffff32e033107110d7303fa2fcf4eb11ee1cfe54f5017f9a653084c5d1d"
);
assert_eq!(sealed.container.len(), 125);
assert_eq!(sealed.wrapped_dek.len(), WRAPPED_DEK_BYTES);
}
#[test]
fn the_header_is_associated_data_so_editing_it_fails_the_open() {
let sealed = seal_fixture(b"payload");
for (label, at) in [
("version", 9),
("digest", 12),
("byte_size", 44),
("stream_nonce", 52),
("media_type", 79),
("kek_id", 92),
] {
let mut tampered = sealed.container.clone();
tampered[at] ^= 0xff;
let err =
open(&tampered, &KEK, &sealed.wrap_nonce, &sealed.wrapped_dek).expect_err(label);
assert!(matches!(err, BlobError::Integrity(_)), "{label}: {err:?}");
}
}
#[test]
fn a_flipped_ciphertext_byte_fails_authentication() {
let sealed = seal_fixture(b"payload");
let mut tampered = sealed.container.clone();
let last = tampered.len() - 1;
tampered[last] ^= 0x01;
assert!(matches!(
open(&tampered, &KEK, &sealed.wrap_nonce, &sealed.wrapped_dek),
Err(BlobError::Integrity(_))
));
}
#[test]
fn truncation_is_detected_because_the_last_chunk_is_sealed_as_last() {
let plaintext = vec![0x5a; BLOB_CHUNK_BYTES + 32];
let sealed = seal_fixture(&plaintext);
let (_, header_len) = Header::decode(&sealed.container).expect("header");
let (first, after_first) = read_chunk(&sealed.container, header_len).expect("chunk");
assert!(after_first < sealed.container.len(), "expected two chunks");
let mut truncated = sealed.container[..header_len].to_vec();
push_chunk(&mut truncated, first).expect("frame");
let err = open(&truncated, &KEK, &sealed.wrap_nonce, &sealed.wrapped_dek)
.expect_err("truncated container");
assert!(matches!(err, BlobError::Integrity(_)), "{err:?}");
}
#[test]
fn the_wrong_kek_cannot_open_it_and_says_nothing_about_why() {
let sealed = seal_fixture(b"secret");
let wrong = [0x99; DEK_BYTES];
let err = open(
&sealed.container,
&wrong,
&sealed.wrap_nonce,
&sealed.wrapped_dek,
)
.expect_err("wrong kek");
let BlobError::Integrity(message) = err else {
panic!("expected an integrity failure");
};
assert!(message.contains("wrong kek or altered header"), "{message}");
}
#[test]
fn rewrapping_changes_only_the_key_not_the_ciphertext() {
let sealed = seal_fixture(b"rotate me");
let new_kek = [0x77; DEK_BYTES];
let new_nonce = [0x88; WRAP_NONCE_BYTES];
let rewrapped = rewrap(
&sealed.container,
&KEK,
&new_kek,
&sealed.wrap_nonce,
&sealed.wrapped_dek,
&new_nonce,
)
.expect("rewrap");
assert_ne!(rewrapped, sealed.wrapped_dek);
let (_, opened) =
open(&sealed.container, &new_kek, &new_nonce, &rewrapped).expect("open rewrapped");
assert_eq!(opened, b"rotate me");
assert!(open(&sealed.container, &KEK, &new_nonce, &rewrapped).is_err());
}
#[test]
fn a_generated_seal_uses_fresh_material_every_time() {
let a = seal(b"same bytes", "text/plain", &KEK, "kek-test").expect("seal");
let b = seal(b"same bytes", "text/plain", &KEK, "kek-test").expect("seal");
assert_eq!(a.digest_hex, b.digest_hex);
assert_ne!(a.container, b.container);
assert_ne!(a.wrap_nonce, b.wrap_nonce);
let (_, opened) = open(&a.container, &KEK, &a.wrap_nonce, &a.wrapped_dek).expect("open");
assert_eq!(opened, b"same bytes");
}
#[test]
fn a_chunk_can_be_opened_where_it_lies_without_reading_the_ones_before_it() {
let plaintext: Vec<u8> = (0..BLOB_CHUNK_BYTES * 2 + 7)
.map(|i| (i % 251) as u8)
.collect();
let sealed = seal_fixture(&plaintext);
let (header, header_len) = Header::decode(&sealed.container).expect("header");
let aad = &sealed.container[..header_len];
let dek = unwrap_dek(&sealed.wrapped_dek, &KEK, &sealed.wrap_nonce, aad).expect("unwrap");
let count = chunk_count(header.byte_size);
assert_eq!(count, 3);
let mut reassembled = Vec::new();
for index in 0..count {
let at = chunk_offset(header_len, index) as usize;
let (chunk, _) = read_chunk(&sealed.container, at).expect("frame");
let part = open_chunk(
aad,
&dek,
&header.stream_nonce,
index as u32,
index == count - 1,
chunk,
)
.expect("chunk opens on its own");
reassembled.extend_from_slice(&part);
}
assert_eq!(reassembled, plaintext);
let (middle, _) =
read_chunk(&sealed.container, chunk_offset(header_len, 0) as usize).expect("frame");
assert!(
open_chunk(aad, &dek, &header.stream_nonce, 0, true, middle).is_err(),
"chunk 0 must not open as the final chunk"
);
let (final_chunk, _) =
read_chunk(&sealed.container, chunk_offset(header_len, 2) as usize).expect("frame");
assert!(
open_chunk(aad, &dek, &header.stream_nonce, 2, false, final_chunk).is_err(),
"the final chunk must not open as a middle one"
);
assert!(
open_chunk(aad, &dek, &header.stream_nonce, 1, false, middle).is_err(),
"chunk 0 must not open at position 1"
);
}
#[test]
fn a_container_that_is_not_one_is_refused_before_any_crypto() {
for (label, bytes) in [
("empty", Vec::new()),
("short", vec![0u8; 10]),
("bad magic", {
let mut v = seal_fixture(b"x").container;
v[0] = b'X';
v
}),
("future version", {
let mut v = seal_fixture(b"x").container;
v[9] = 0x02;
v
}),
] {
assert!(
matches!(Header::decode(&bytes), Err(BlobError::Integrity(_))),
"{label} decoded"
);
}
}
}