use tink_core::{utils::wrap_err, TinkError};
pub fn new(h: &tink_core::keyset::Handle) -> Result<Box<dyn tink_core::Aead>, TinkError> {
new_with_key_manager(h, None)
}
fn new_with_key_manager(
h: &tink_core::keyset::Handle,
km: Option<std::sync::Arc<dyn tink_core::registry::KeyManager>>,
) -> Result<Box<dyn tink_core::Aead>, TinkError> {
let ps = h
.primitives_with_key_manager(km)
.map_err(|e| wrap_err("aead::factory: cannot obtain primitive set", e))?;
let ret = WrappedAead::new(ps)?;
Ok(Box::new(ret))
}
#[derive(Clone)]
struct WrappedAead {
ps: tink_core::primitiveset::TypedPrimitiveSet<Box<dyn tink_core::Aead>>,
}
impl WrappedAead {
fn new(ps: tink_core::primitiveset::PrimitiveSet) -> Result<WrappedAead, TinkError> {
let entry = match &ps.primary {
None => return Err("aead::factory: no primary primitive".into()),
Some(p) => p,
};
match entry.primitive {
tink_core::Primitive::Aead(_) => {}
_ => return Err("aead::factory: not an AEAD primitive".into()),
};
for (_, primitives) in ps.entries.iter() {
for p in primitives {
match p.primitive {
tink_core::Primitive::Aead(_) => {}
_ => return Err("aead::factory: not an AEAD primitive".into()),
};
}
}
Ok(WrappedAead { ps: ps.into() })
}
}
impl tink_core::Aead for WrappedAead {
fn encrypt(&self, pt: &[u8], aad: &[u8]) -> Result<Vec<u8>, TinkError> {
let primary = self
.ps
.primary
.as_ref()
.ok_or_else(|| TinkError::new("no primary"))?;
let ct = primary.primitive.encrypt(pt, aad)?;
let mut ret = Vec::with_capacity(primary.prefix.len() + ct.len());
ret.extend_from_slice(&primary.prefix);
ret.extend_from_slice(&ct);
Ok(ret)
}
fn decrypt(&self, ct: &[u8], aad: &[u8]) -> Result<Vec<u8>, TinkError> {
let prefix_size = tink_core::cryptofmt::NON_RAW_PREFIX_SIZE;
if ct.len() > prefix_size {
let prefix = &ct[..prefix_size];
let ct_no_prefix = &ct[prefix_size..];
if let Some(entries) = self.ps.entries_for_prefix(prefix) {
for entry in entries {
if let Ok(pt) = entry.primitive.decrypt(ct_no_prefix, aad) {
return Ok(pt);
}
}
}
}
if let Some(entries) = self.ps.raw_entries() {
for entry in entries {
if let Ok(pt) = entry.primitive.decrypt(ct, aad) {
return Ok(pt);
}
}
}
Err("aead::decrypt: decryption failed".into())
}
}