use serde::{Deserialize, Serialize};
use std::io::Read;
use tls_codec::{
Deserialize as TlsDeserializeTrait, DeserializeBytes as TlsDeserializeBytesTrait,
Serialize as TlsSerializeTrait, TlsDeserialize, TlsDeserializeBytes, TlsSerialize, TlsSize,
VLBytes,
};
use crate::component::ComponentId;
#[derive(thiserror::Error, Debug, PartialEq, Eq, Clone)]
pub enum SafeAadError {
#[error("duplicate component id in SafeAAD: {0}")]
DuplicateComponentId(ComponentId),
#[error("SafeAAD items are not sorted by component id in increasing order")]
ItemsNotSortedAscending,
#[error("codec error: {0}")]
Codec(String),
}
#[derive(
Clone,
Debug,
PartialEq,
Eq,
Serialize,
Deserialize,
TlsSerialize,
TlsDeserialize,
TlsDeserializeBytes,
TlsSize,
)]
pub struct SafeAadItem {
component_id: ComponentId,
aad_item_data: VLBytes,
}
impl SafeAadItem {
pub fn new(component_id: ComponentId, data: Vec<u8>) -> Self {
Self {
component_id,
aad_item_data: data.into(),
}
}
pub fn component_id(&self) -> ComponentId {
self.component_id
}
pub fn data(&self) -> &[u8] {
self.aad_item_data.as_slice()
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize, TlsSerialize, TlsSize)]
pub struct SafeAad {
aad_items: Vec<SafeAadItem>,
}
impl SafeAad {
pub fn from_items(items: Vec<SafeAadItem>) -> Result<Self, SafeAadError> {
Self::validate(&items)?;
Ok(Self { aad_items: items })
}
pub fn empty() -> Self {
Self {
aad_items: Vec::new(),
}
}
pub fn items(&self) -> &[SafeAadItem] {
&self.aad_items
}
pub fn get(&self, component_id: ComponentId) -> Option<&[u8]> {
self.aad_items
.binary_search_by_key(&component_id, SafeAadItem::component_id)
.ok()
.map(|index| self.aad_items[index].data())
}
pub fn is_empty(&self) -> bool {
self.aad_items.is_empty()
}
pub fn len(&self) -> usize {
self.aad_items.len()
}
fn validate(items: &[SafeAadItem]) -> Result<(), SafeAadError> {
let mut previous: Option<ComponentId> = None;
for item in items {
if let Some(prev) = previous {
if item.component_id == prev {
return Err(SafeAadError::DuplicateComponentId(item.component_id));
}
if item.component_id < prev {
return Err(SafeAadError::ItemsNotSortedAscending);
}
}
previous = Some(item.component_id);
}
Ok(())
}
}
impl TlsDeserializeTrait for SafeAad {
fn tls_deserialize<R: Read>(bytes: &mut R) -> Result<Self, tls_codec::Error> {
let aad_items = Vec::<SafeAadItem>::tls_deserialize(bytes)?;
SafeAad::from_items(aad_items)
.map_err(|err| tls_codec::Error::DecodingError(err.to_string()))
}
}
impl TlsDeserializeBytesTrait for SafeAad {
fn tls_deserialize_bytes(bytes: &[u8]) -> Result<(Self, &[u8]), tls_codec::Error> {
let (aad_items, rest) = Vec::<SafeAadItem>::tls_deserialize_bytes(bytes)?;
let aad = SafeAad::from_items(aad_items)
.map_err(|err| tls_codec::Error::DecodingError(err.to_string()))?;
Ok((aad, rest))
}
}
pub(crate) fn assemble_authenticated_data(
safe_aad: &SafeAad,
tail: &[u8],
) -> Result<Vec<u8>, SafeAadError> {
let mut out = safe_aad
.tls_serialize_detached()
.map_err(|err| SafeAadError::Codec(err.to_string()))?;
out.extend_from_slice(tail);
Ok(out)
}
pub(crate) fn parse_authenticated_data_prefix(
bytes: &[u8],
) -> Result<(SafeAad, usize), SafeAadError> {
let (parsed, remainder) = SafeAad::tls_deserialize_bytes(bytes)
.map_err(|err| SafeAadError::Codec(err.to_string()))?;
let prefix_len = bytes.len() - remainder.len();
Ok((parsed, prefix_len))
}
#[cfg(test)]
mod tests {
use super::*;
use tls_codec::{Deserialize, Serialize};
fn item(id: ComponentId, data: &[u8]) -> SafeAadItem {
SafeAadItem::new(id, data.to_vec())
}
#[test]
fn roundtrip_non_empty() {
let safe_aad = SafeAad::from_items(vec![
item(1, b"first"),
item(7, b""),
item(42, b"last item bytes"),
])
.unwrap();
let bytes = safe_aad.tls_serialize_detached().unwrap();
let parsed = SafeAad::tls_deserialize_exact(&bytes).unwrap();
assert_eq!(parsed, safe_aad);
let reserialized = parsed.tls_serialize_detached().unwrap();
assert_eq!(reserialized, bytes);
}
#[test]
fn empty_is_length_prefix_only() {
let safe_aad = SafeAad::empty();
let bytes = safe_aad.tls_serialize_detached().unwrap();
assert_eq!(bytes, vec![0x00]);
let parsed = SafeAad::tls_deserialize_exact(&bytes).unwrap();
assert!(parsed.is_empty());
}
#[test]
fn from_items_rejects_duplicates() {
let err = SafeAad::from_items(vec![item(3, b"a"), item(3, b"b")]).unwrap_err();
assert_eq!(err, SafeAadError::DuplicateComponentId(3));
}
#[test]
fn from_items_rejects_misordered() {
let err = SafeAad::from_items(vec![item(9, b""), item(2, b"")]).unwrap_err();
assert_eq!(err, SafeAadError::ItemsNotSortedAscending);
}
#[test]
fn deserialize_rejects_misordered() {
let raw_items: Vec<SafeAadItem> = vec![item(5, b"x"), item(1, b"y")];
let raw_bytes = raw_items.tls_serialize_detached().unwrap();
let err = SafeAad::tls_deserialize_exact(&raw_bytes).unwrap_err();
match err {
tls_codec::Error::DecodingError(message) => {
assert!(
message.contains("not sorted"),
"unexpected error message: {message}"
);
}
other => panic!("unexpected error variant: {other:?}"),
}
}
#[test]
fn deserialize_rejects_duplicates() {
let raw_items: Vec<SafeAadItem> = vec![item(4, b""), item(4, b"")];
let raw_bytes = raw_items.tls_serialize_detached().unwrap();
let err = SafeAad::tls_deserialize_exact(&raw_bytes).unwrap_err();
match err {
tls_codec::Error::DecodingError(message) => {
assert!(
message.contains("duplicate"),
"unexpected error message: {message}"
);
}
other => panic!("unexpected error variant: {other:?}"),
}
}
#[test]
fn boundary_component_ids() {
let safe_aad = SafeAad::from_items(vec![item(0, b"min"), item(u16::MAX, b"max")]).unwrap();
let bytes = safe_aad.tls_serialize_detached().unwrap();
let parsed = SafeAad::tls_deserialize_exact(&bytes).unwrap();
assert_eq!(parsed.get(0), Some(b"min".as_slice()));
assert_eq!(parsed.get(u16::MAX), Some(b"max".as_slice()));
}
#[test]
fn get_returns_none_for_missing() {
let safe_aad = SafeAad::from_items(vec![item(1, b"a"), item(10, b"b")]).unwrap();
assert_eq!(safe_aad.get(5), None);
assert_eq!(safe_aad.get(1), Some(b"a".as_slice()));
assert_eq!(safe_aad.get(10), Some(b"b".as_slice()));
}
#[test]
fn assemble_and_parse_authenticated_data_roundtrip() {
let safe_aad =
SafeAad::from_items(vec![item(2, b"safe-aad-data"), item(8, b"more")]).unwrap();
let tail = b"caller tail bytes";
let combined = assemble_authenticated_data(&safe_aad, tail).unwrap();
let (parsed, prefix_len) = parse_authenticated_data_prefix(&combined).unwrap();
assert_eq!(parsed, safe_aad);
assert_eq!(&combined[prefix_len..], tail);
}
}