use zeroize::Zeroizing;
use crate::primitives::domain::DS_ITEM;
use crate::primitives::{Aead, Kdf, PrimitiveSuite};
use crate::{Error, Result};
pub const RECORD_SUITE_XCHACHA20POLY1305: u8 = 0x01;
pub struct SealCtx<'a> {
pub domain: &'a str,
pub vault: &'a [u8],
pub id: &'a [u8],
pub version: &'a [u8],
}
fn push_lp(out: &mut Vec<u8>, field: &[u8], what: &'static str) -> Result<()> {
let len = u32::try_from(field.len()).map_err(|_| Error::Malformed(what))?;
out.extend_from_slice(&len.to_be_bytes());
out.extend_from_slice(field);
Ok(())
}
pub fn record_aad(suite: u8, ctx: &SealCtx<'_>) -> Result<Vec<u8>> {
let domain = ctx.domain.as_bytes();
let mut aad = Vec::with_capacity(
1 + 4
+ DS_ITEM.len()
+ 4
+ domain.len()
+ 4
+ ctx.vault.len()
+ 4
+ ctx.id.len()
+ 4
+ ctx.version.len(),
);
aad.push(suite);
push_lp(&mut aad, DS_ITEM, "record AAD: label too long")?;
push_lp(&mut aad, domain, "record AAD: domain too long")?;
push_lp(&mut aad, ctx.vault, "record AAD: vault id too long")?;
push_lp(&mut aad, ctx.id, "record AAD: record id too long")?;
push_lp(&mut aad, ctx.version, "record AAD: version too long")?;
Ok(aad)
}
fn derive_item_key<S: PrimitiveSuite>(k: &[u8]) -> Result<Zeroizing<[u8; 32]>> {
Ok(Zeroizing::new(S::Kdf::derive_32(k, &[], DS_ITEM)?))
}
pub fn seal_record<S: PrimitiveSuite>(
k: &[u8],
ctx: &SealCtx<'_>,
plaintext: &[u8],
) -> Result<Vec<u8>> {
let k_aead = derive_item_key::<S>(k)?;
let aad = record_aad(RECORD_SUITE_XCHACHA20POLY1305, ctx)?;
let mut sealed = S::Aead::seal(&k_aead[..], plaintext, &aad)?;
let mut out = Vec::with_capacity(1 + sealed.len());
out.push(RECORD_SUITE_XCHACHA20POLY1305);
out.append(&mut sealed);
Ok(out)
}
pub fn unseal_record<S: PrimitiveSuite>(
k: &[u8],
ctx: &SealCtx<'_>,
sealed: &[u8],
) -> Result<Vec<u8>> {
let (&suite, body) = sealed
.split_first()
.ok_or(Error::Malformed("record: empty sealed blob"))?;
if suite != RECORD_SUITE_XCHACHA20POLY1305 {
return Err(Error::Malformed("record: unknown suite tag"));
}
let k_aead = derive_item_key::<S>(k)?;
let aad = record_aad(suite, ctx)?;
S::Aead::open(&k_aead[..], body, &aad)
}
#[cfg(all(test, feature = "std-primitives"))]
mod tests {
use super::*;
use crate::primitives::{Aead, ChaCha20Poly1305, StdPrimitives};
fn ctx<'a>(domain: &'a str, vault: &'a [u8], id: &'a [u8], version: &'a [u8]) -> SealCtx<'a> {
SealCtx {
domain,
vault,
id,
version,
}
}
fn hex(b: &[u8]) -> String {
b.iter().map(|x| format!("{:02x}", x)).collect()
}
#[test]
fn roundtrip() {
let k = [0x42u8; 32];
let c = ctx("item", b"vault-1", b"id-abc", &[7]);
let pt = b"a connect token";
let sealed = seal_record::<StdPrimitives>(&k, &c, pt).unwrap();
assert_eq!(sealed[0], RECORD_SUITE_XCHACHA20POLY1305);
let opened = unseal_record::<StdPrimitives>(&k, &c, &sealed).unwrap();
assert_eq!(opened, pt);
}
#[test]
fn empty_plaintext_roundtrips() {
let k = [0x42u8; 32];
let c = ctx("item", b"v", b"i", &[0]);
let sealed = seal_record::<StdPrimitives>(&k, &c, b"").unwrap();
assert_eq!(
unseal_record::<StdPrimitives>(&k, &c, &sealed).unwrap(),
b""
);
}
#[test]
fn empty_version_roundtrips() {
let k = [0x42u8; 32];
let c = ctx("item", b"v", b"id", b"");
let sealed = seal_record::<StdPrimitives>(&k, &c, b"x").unwrap();
assert_eq!(
unseal_record::<StdPrimitives>(&k, &c, &sealed).unwrap(),
b"x"
);
}
#[test]
fn wrong_vault_fails() {
let k = [0x42u8; 32];
let sealed =
seal_record::<StdPrimitives>(&k, &ctx("item", b"vault-1", b"id", &[1]), b"x").unwrap();
let bad =
unseal_record::<StdPrimitives>(&k, &ctx("item", b"vault-2", b"id", &[1]), &sealed);
assert!(matches!(bad, Err(Error::SealDecryptionFailed)));
}
#[test]
fn wrong_id_fails() {
let k = [0x42u8; 32];
let sealed =
seal_record::<StdPrimitives>(&k, &ctx("item", b"v", b"id-A", &[1]), b"x").unwrap();
let bad = unseal_record::<StdPrimitives>(&k, &ctx("item", b"v", b"id-B", &[1]), &sealed);
assert!(matches!(bad, Err(Error::SealDecryptionFailed)));
}
#[test]
fn wrong_version_fails() {
let k = [0x42u8; 32];
let sealed =
seal_record::<StdPrimitives>(&k, &ctx("item", b"v", b"id", &[1]), b"x").unwrap();
let bad = unseal_record::<StdPrimitives>(&k, &ctx("item", b"v", b"id", &[2]), &sealed);
assert!(matches!(bad, Err(Error::SealDecryptionFailed)));
}
#[test]
fn version_length_extension_fails() {
let k = [0x42u8; 32];
let sealed =
seal_record::<StdPrimitives>(&k, &ctx("item", b"v", b"id", &[1]), b"x").unwrap();
let bad = unseal_record::<StdPrimitives>(&k, &ctx("item", b"v", b"id", &[0, 1]), &sealed);
assert!(matches!(bad, Err(Error::SealDecryptionFailed)));
}
#[test]
fn wrong_domain_fails() {
let k = [0x42u8; 32];
let sealed =
seal_record::<StdPrimitives>(&k, &ctx("item", b"v", b"id", &[1]), b"x").unwrap();
let bad = unseal_record::<StdPrimitives>(&k, &ctx("keyset", b"v", b"id", &[1]), &sealed);
assert!(matches!(bad, Err(Error::SealDecryptionFailed)));
}
#[test]
fn wrong_key_fails() {
let sealed =
seal_record::<StdPrimitives>(&[1u8; 32], &ctx("item", b"v", b"id", &[1]), b"x")
.unwrap();
let bad =
unseal_record::<StdPrimitives>(&[2u8; 32], &ctx("item", b"v", b"id", &[1]), &sealed);
assert!(matches!(bad, Err(Error::SealDecryptionFailed)));
}
#[test]
fn empty_blob_is_malformed() {
let bad = unseal_record::<StdPrimitives>(&[0u8; 32], &ctx("item", b"v", b"id", &[1]), b"");
assert!(matches!(bad, Err(Error::Malformed(_))));
}
#[test]
fn unknown_suite_is_malformed() {
let k = [0x42u8; 32];
let mut sealed =
seal_record::<StdPrimitives>(&k, &ctx("item", b"v", b"id", &[1]), b"x").unwrap();
sealed[0] = 0x02; let bad = unseal_record::<StdPrimitives>(&k, &ctx("item", b"v", b"id", &[1]), &sealed);
assert!(matches!(bad, Err(Error::Malformed(_))));
}
#[test]
fn lp_ambiguity_is_resolved() {
let a = record_aad(0x01, &ctx("item", b"ab", b"c", &[1])).unwrap();
let b = record_aad(0x01, &ctx("item", b"a", b"bc", &[1])).unwrap();
assert_ne!(a, b);
}
const CV_K: [u8; 32] = [0x11u8; 32];
const CV_NONCE: [u8; 24] = [0x22u8; 24];
const CV_DOMAIN: &str = "item";
const CV_VAULT: &[u8] = b"vault-7";
const CV_ID: &[u8] = &[0xAA, 0xBB, 0xCC, 0xDD];
const CV_VERSION: &[u8] = &[0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08];
const CV_PT: &[u8] = b"the lazy dog jumps over...";
fn cv_ctx() -> SealCtx<'static> {
ctx(CV_DOMAIN, CV_VAULT, CV_ID, CV_VERSION)
}
#[test]
fn conformance_record_aad() {
let aad = record_aad(RECORD_SUITE_XCHACHA20POLY1305, &cv_ctx()).unwrap();
assert_eq!(
hex(&aad),
"010000000c737564702f76312f6974656d000000046974656d000000077661756c742d3700000004aabbccdd000000080102030405060708"
);
}
#[test]
fn conformance_item_key() {
let k_aead = derive_item_key::<StdPrimitives>(&CV_K).unwrap();
assert_eq!(
hex(&k_aead[..]),
"d9e525d7f8047ad0c47bc270f44e22a7a4038d2fb7df863924128481efe83823"
);
}
#[test]
fn conformance_sealed_fixed_nonce() {
let k_aead = derive_item_key::<StdPrimitives>(&CV_K).unwrap();
let aad = record_aad(RECORD_SUITE_XCHACHA20POLY1305, &cv_ctx()).unwrap();
let ct = ChaCha20Poly1305::encrypt(&k_aead[..], &CV_NONCE, CV_PT, &aad).unwrap();
let mut sealed = Vec::new();
sealed.push(RECORD_SUITE_XCHACHA20POLY1305);
sealed.extend_from_slice(&CV_NONCE);
sealed.extend_from_slice(&ct);
assert_eq!(
hex(&sealed),
"0122222222222222222222222222222222222222222222222291131d0f0ef48770f42cb1bd5ef3915479ad080de28b148392796ccd6f88a3eeb1c5fe3a3bff54a793be"
);
let opened = unseal_record::<StdPrimitives>(&CV_K, &cv_ctx(), &sealed).unwrap();
assert_eq!(opened, CV_PT);
}
}