extern crate alloc;
use alloc::vec::Vec;
use crate::kairos::{KAIROS_WIRE_LEN, Kairos};
use crate::metis::Cut;
use crate::metis::dot_set::HAVE_SET_WIRE_HEADER_LEN;
use thiserror::Error;
use super::super::Declaration;
use super::super::arrival::Admission;
use super::super::gate::{EpochAddress, InvalidEpochAddress};
use super::error::WireCutError;
use super::shared::{
WIRE_ADDRESS_LEN, WIRE_DOT_LEN, cut_embedding_len, decode_cut, encode_cut_into,
encode_dot_into, read_address_parts, read_u32,
};
use alloc::collections::BTreeSet;
const DECLARATION_WIRE_V1: u8 = 0x01;
const DECLARATION_WIRE_V2: u8 = 0x02;
const DECLARATION_WIRE_PREFIX_LEN: usize = 1 + 8 + WIRE_DOT_LEN + KAIROS_WIRE_LEN;
const DECLARATION_WIRE_MIN_LEN: usize = DECLARATION_WIRE_PREFIX_LEN + HAVE_SET_WIRE_HEADER_LEN;
#[non_exhaustive]
#[derive(Error, Debug, Clone, Copy, PartialEq, Eq)]
pub enum DeclarationDecodeError {
#[error("unknown declaration wire version: {0:#04x}")]
UnknownVersion(u8),
#[error("unexpected declaration frame length: expected {expected}, found {found}")]
UnexpectedLength {
expected: usize,
found: usize,
},
#[error("declaration address: {0}")]
Address(#[from] InvalidEpochAddress),
#[error("declaration rank: {0}")]
Rank(crate::kairos::DecodeError),
#[error("declaration cut: {0}")]
Cut(#[from] WireCutError),
#[error("declaration admission carries no joiners")]
EmptyAdmission,
#[error("declaration joiners not ascending: {found} after {previous}")]
NonAscendingJoiners {
previous: u32,
found: u32,
},
#[error("declaration admission base at generation {base}, declaration at {declaration}")]
AdmissionBaseMismatch {
base: u64,
declaration: u64,
},
}
impl Declaration {
#[must_use]
pub fn to_bytes(&self) -> Vec<u8> {
let mut out = Vec::with_capacity(self.encoded_len());
out.push(if self.admission.is_some() {
DECLARATION_WIRE_V2
} else {
DECLARATION_WIRE_V1
});
out.extend_from_slice(&self.address.generation().to_be_bytes());
encode_dot_into(&mut out, self.address.declaration());
out.extend_from_slice(&self.rank.to_bytes());
encode_cut_into(&mut out, self.cut.as_vector());
if let Some(admission) = &self.admission {
out.extend_from_slice(&admission.base().generation().to_be_bytes());
encode_dot_into(&mut out, admission.base().declaration());
let joiner_count = u32::try_from(admission.joiners().count()).unwrap_or(u32::MAX);
out.extend_from_slice(&joiner_count.to_be_bytes());
for joiner in admission.joiners() {
out.extend_from_slice(&joiner.to_be_bytes());
}
if let Some(commitment) = &self.admission_commitment {
out.extend_from_slice(commitment);
}
}
out
}
#[must_use]
pub fn encoded_len(&self) -> usize {
DECLARATION_WIRE_PREFIX_LEN
+ cut_embedding_len(self.cut.as_vector())
+ self.admission.as_ref().map_or(0, |admission| {
WIRE_ADDRESS_LEN + 4 + admission.joiners().count() * 4 + 32
})
}
pub fn from_bytes(bytes: &[u8]) -> Result<Self, DeclarationDecodeError> {
if bytes.len() < DECLARATION_WIRE_MIN_LEN {
return Err(DeclarationDecodeError::UnexpectedLength {
expected: DECLARATION_WIRE_MIN_LEN,
found: bytes.len(),
});
}
let version = bytes[0];
if version != DECLARATION_WIRE_V1 && version != DECLARATION_WIRE_V2 {
return Err(DeclarationDecodeError::UnknownVersion(version));
}
let short = DeclarationDecodeError::UnexpectedLength {
expected: DECLARATION_WIRE_MIN_LEN,
found: bytes.len(),
};
let ((generation, dot), rest) = read_address_parts(&bytes[1..]).ok_or(short)?;
let address = EpochAddress::try_from_parts(generation, dot)?;
let stamp = rest.get(..KAIROS_WIRE_LEN).ok_or(short)?;
let rank = Kairos::from_bytes(stamp).map_err(DeclarationDecodeError::Rank)?;
let rest = &rest[KAIROS_WIRE_LEN..];
let (cut, tail) = decode_cut(rest)?;
let (admission, admission_commitment, tail) = if version == DECLARATION_WIRE_V2 {
let (admission, rest) = read_admission(bytes, tail, address.generation())?;
let commitment: [u8; 32] = rest
.get(..32)
.and_then(|bytes| bytes.try_into().ok())
.ok_or(short)?;
(Some(admission), Some(commitment), &rest[32..])
} else {
(None, None, tail)
};
if !tail.is_empty() {
return Err(DeclarationDecodeError::UnexpectedLength {
expected: bytes.len() - tail.len(),
found: bytes.len(),
});
}
Ok(Self {
address,
rank,
cut: Cut::from_witnessed(cut),
admission,
admission_commitment,
})
}
}
fn read_admission<'a>(
bytes: &[u8],
rest: &'a [u8],
declared: u64,
) -> Result<(Admission, &'a [u8]), DeclarationDecodeError> {
let short = DeclarationDecodeError::UnexpectedLength {
expected: bytes.len() + 1,
found: bytes.len(),
};
let ((base_generation, base_dot), rest) = read_address_parts(rest).ok_or(short)?;
let base = EpochAddress::try_from_parts(base_generation, base_dot)?;
if base_generation.checked_add(1) != Some(declared) {
return Err(DeclarationDecodeError::AdmissionBaseMismatch {
base: base_generation,
declaration: declared,
});
}
let (count, mut rest) = read_u32(rest).ok_or(short)?;
if count == 0 {
return Err(DeclarationDecodeError::EmptyAdmission);
}
if (count as usize).saturating_mul(4) > rest.len() {
return Err(short);
}
let mut joiners = BTreeSet::new();
let mut previous: Option<u32> = None;
for _ in 0..count {
let (joiner, tail) = read_u32(rest).ok_or(short)?;
if previous.is_some_and(|previous| joiner <= previous) {
return Err(DeclarationDecodeError::NonAscendingJoiners {
previous: previous.unwrap_or(0),
found: joiner,
});
}
previous = Some(joiner);
let _ = joiners.insert(joiner);
rest = tail;
}
let admission =
Admission::new(base, joiners).map_err(|_refusal| DeclarationDecodeError::EmptyAdmission)?;
Ok((admission, rest))
}