use crate::{Nonce, Unspecified};
use bytes::{Buf, Bytes};
use std::array;
use vitaminc_protected::{Controlled, Protected};
pub struct CipherTextReader(Bytes);
impl CipherTextReader {
pub(super) fn new(bytes: Bytes) -> Self {
Self(bytes)
}
pub fn read_version(self) -> Result<VersionedReader, Unspecified> {
let mut buf = self.0;
if buf.is_empty() {
return Err(Unspecified);
}
let version = buf.get_u8();
if version != super::WIRE_VERSION {
return Err(Unspecified);
}
Ok(VersionedReader(buf))
}
}
pub struct VersionedReader(Bytes);
impl VersionedReader {
pub fn read_nonce<const N: usize>(
self,
) -> Result<(Nonce<N>, CiphertextAndTagReader), Unspecified> {
if self.0.len() < N {
return Err(Unspecified);
}
let mut buf = self.0.take(N);
let nonce_inner: [u8; N] = array::from_fn(|_| buf.get_u8());
Ok((
Nonce::new(nonce_inner),
CiphertextAndTagReader::new(buf.into_inner()),
))
}
}
pub struct CiphertextAndTagReader(Protected<Vec<u8>>);
impl CiphertextAndTagReader {
fn new(bytes: Bytes) -> Self {
Self(Protected::new(bytes.into()))
}
pub fn accepts_plaintext_ok<E>(
self,
f: impl FnOnce(&mut [u8]) -> Result<usize, E>,
) -> Plaintext<E> {
Plaintext(self.0.map_ok(|mut raw| {
let len = f(&mut raw)?;
raw.truncate(len);
Ok(raw)
}))
}
}
pub struct Plaintext<E>(Result<Protected<Vec<u8>>, E>);
impl<E> Plaintext<E> {
pub fn read(self) -> Result<Protected<Vec<u8>>, E> {
self.0
}
}
#[cfg(test)]
mod tests {
use super::{CipherTextReader, VersionedReader};
use crate::ciphertext::WIRE_VERSION;
use bytes::Bytes;
use vitaminc_protected::Controlled;
const N: usize = 4;
fn versioned(bytes: &'static [u8]) -> VersionedReader {
let mut buf = vec![WIRE_VERSION];
buf.extend_from_slice(bytes);
CipherTextReader::new(Bytes::from(buf))
.read_version()
.expect("known version")
}
#[test]
fn read_version_empty_buffer_errors() {
let reader = CipherTextReader::new(Bytes::new());
assert!(reader.read_version().is_err());
}
#[test]
fn read_version_unknown_version_errors() {
for version in 0..=u8::MAX {
if version == WIRE_VERSION {
continue;
}
let reader = CipherTextReader::new(Bytes::from(vec![version, 1, 2, 3, 4]));
assert!(
reader.read_version().is_err(),
"version {version} must be rejected"
);
}
}
#[test]
fn read_version_strips_exactly_one_byte() {
let (nonce, rest) = versioned(&[1, 2, 3, 4, 5])
.read_nonce::<N>()
.expect("nonce available");
assert_eq!(nonce.into_inner(), [1, 2, 3, 4]);
assert_eq!(rest.0.risky_ref(), &vec![5]);
}
#[test]
fn read_nonce_fewer_than_n_bytes_errors() {
assert!(versioned(&[1, 2, 3]).read_nonce::<N>().is_err());
}
#[test]
fn read_nonce_exactly_n_bytes_succeeds() {
let (nonce, rest) = versioned(&[1, 2, 3, 4])
.read_nonce::<N>()
.expect("N bytes available");
assert_eq!(nonce.into_inner(), [1, 2, 3, 4]);
assert!(rest.0.risky_ref().is_empty());
}
#[test]
fn read_nonce_more_than_n_bytes_succeeds() {
let (nonce, rest) = versioned(&[1, 2, 3, 4, 5, 6])
.read_nonce::<N>()
.expect("more than N bytes available");
assert_eq!(nonce.into_inner(), [1, 2, 3, 4]);
assert_eq!(rest.0.risky_ref(), &vec![5, 6]);
}
}