extern crate alloc;
use alloc::collections::BTreeSet;
use alloc::vec::Vec;
use crate::metis::VersionVector;
use crate::metis::dot::Dot;
use crate::metis::dot_set::HAVE_SET_WIRE_HEADER_LEN;
use thiserror::Error;
use super::super::SealedEpoch;
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_address_into,
encode_cut_into, encode_dot_into, read_address_parts, read_dot, read_u32,
};
const SEAL_WIRE_V1: u8 = 0x01;
const SEAL_WIRE_V2: u8 = 0x02;
pub(super) const SEAL_WIRE_MIN_LEN: usize = 1 + WIRE_ADDRESS_LEN + 4 + 4 + HAVE_SET_WIRE_HEADER_LEN;
const SEAL_DEFAULT_MAX_JOINERS: usize = 1024;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct SealDecodeBudget {
max_candidates: usize,
max_joiners: usize,
}
impl SealDecodeBudget {
#[must_use]
pub const fn new(max_candidates: usize) -> Self {
Self {
max_candidates,
max_joiners: SEAL_DEFAULT_MAX_JOINERS,
}
}
#[must_use]
pub const fn with_max_joiners(self, max_joiners: usize) -> Self {
Self {
max_joiners,
..self
}
}
#[must_use]
pub const fn max_candidates(self) -> usize {
self.max_candidates
}
#[must_use]
pub const fn max_joiners(self) -> usize {
self.max_joiners
}
}
#[non_exhaustive]
#[derive(Error, Debug, Clone, Copy, PartialEq, Eq)]
pub enum SealDecodeError {
#[error("unknown seal wire version: {0:#04x}")]
UnknownVersion(u8),
#[error("unexpected seal frame length: expected {expected}, found {found}")]
UnexpectedLength {
expected: usize,
found: usize,
},
#[error("seal address: {0}")]
Address(#[from] InvalidEpochAddress),
#[error("seal candidates not ascending: {found} after {previous}")]
NonAscendingCandidates {
previous: EpochAddress,
found: EpochAddress,
},
#[error("seal ledger dot at station {station} has the non-dot counter zero")]
ZeroLedgerDot {
station: u32,
},
#[error("seal ledger not ascending: {found} after {previous}")]
NonAscendingLedger {
previous: Dot,
found: Dot,
},
#[error("seal frame declares {count} candidates, past the budget's {budget}")]
TooManyCandidates {
count: u64,
budget: u64,
},
#[error("seal frame declares {count} ledger dots, past the budget's {budget}")]
TooManyLedgerDots {
count: u64,
budget: u64,
},
#[error("sealed join: {0}")]
Cut(#[from] WireCutError),
#[error("seal admission carries no joiners")]
EmptyAdmission,
#[error("seal joiners not ascending: {found} after {previous}")]
NonAscendingJoiners {
previous: u32,
found: u32,
},
#[error("seal frame declares {count} joiners, past the budget's {budget}")]
TooManyJoiners {
count: u64,
budget: u64,
},
#[error("seal admission base at generation {base}, record at {record}")]
AdmissionBaseMismatch {
base: u64,
record: u64,
},
}
type DecodedSet<'a, T> = (BTreeSet<T>, &'a [u8]);
fn read_admission<'a>(
bytes: &[u8],
rest: &'a [u8],
declaration: EpochAddress,
budget: Option<SealDecodeBudget>,
) -> Result<(Admission, &'a [u8]), SealDecodeError> {
let short = |rest: &[u8]| SealDecodeError::UnexpectedLength {
expected: bytes.len() + (WIRE_ADDRESS_LEN + 4).saturating_sub(rest.len()),
found: bytes.len(),
};
let ((base_generation, base_dot), rest) =
read_address_parts(rest).ok_or_else(|| short(rest))?;
let base = EpochAddress::try_from_parts(base_generation, base_dot)?;
if base_generation.checked_add(1) != Some(declaration.generation()) {
return Err(SealDecodeError::AdmissionBaseMismatch {
base: base_generation,
record: declaration.generation(),
});
}
let (count, mut rest) = read_u32(rest).ok_or_else(|| short(rest))?;
if count == 0 {
return Err(SealDecodeError::EmptyAdmission);
}
if (count as usize).saturating_mul(4) > rest.len() {
return Err(short(rest));
}
if let Some(budget) = budget
&& count as usize > budget.max_joiners()
{
return Err(SealDecodeError::TooManyJoiners {
count: u64::from(count),
budget: budget.max_joiners() as u64,
});
}
let mut joiners = BTreeSet::new();
let mut previous: Option<u32> = None;
for _ in 0..count {
let (joiner, tail) = read_u32(rest).ok_or_else(|| short(rest))?;
if previous.is_some_and(|previous| joiner <= previous) {
return Err(SealDecodeError::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| SealDecodeError::EmptyAdmission)?;
Ok((admission, rest))
}
fn read_candidate_set<'a>(
bytes: &[u8],
rest: &'a [u8],
budget: Option<SealDecodeBudget>,
) -> Result<DecodedSet<'a, EpochAddress>, SealDecodeError> {
let short = |rest: &[u8], needed: u64| SealDecodeError::UnexpectedLength {
expected: usize::try_from(
(u64::try_from(bytes.len() - rest.len()).unwrap_or(u64::MAX)).saturating_add(needed),
)
.unwrap_or(usize::MAX),
found: bytes.len(),
};
let (count, rest) = read_u32(rest).ok_or_else(|| short(rest, 4))?;
if let Some(budget) = budget
&& u64::from(count) > budget.max_candidates() as u64
{
return Err(SealDecodeError::TooManyCandidates {
count: u64::from(count),
budget: budget.max_candidates() as u64,
});
}
let reserved = 4 + HAVE_SET_WIRE_HEADER_LEN;
let backable = rest.len().saturating_sub(reserved) / WIRE_ADDRESS_LEN;
if count as usize > backable {
let floor = u64::from(count)
.saturating_mul(WIRE_ADDRESS_LEN as u64)
.saturating_add(reserved as u64);
return Err(short(rest, floor));
}
let mut candidates = BTreeSet::new();
let mut previous: Option<EpochAddress> = None;
let mut rest = rest;
for _ in 0..count {
let ((generation, dot), tail) =
read_address_parts(rest).ok_or_else(|| short(rest, WIRE_ADDRESS_LEN as u64))?;
let candidate = EpochAddress::try_from_parts(generation, dot)?;
if let Some(previous) = previous
&& candidate <= previous
{
return Err(SealDecodeError::NonAscendingCandidates {
previous,
found: candidate,
});
}
previous = Some(candidate);
let _ = candidates.insert(candidate);
rest = tail;
}
Ok((candidates, rest))
}
fn read_ledger_set<'a>(
bytes: &[u8],
rest: &'a [u8],
budget: Option<SealDecodeBudget>,
) -> Result<DecodedSet<'a, Dot>, SealDecodeError> {
let short = |rest: &[u8], needed: u64| SealDecodeError::UnexpectedLength {
expected: usize::try_from(
(u64::try_from(bytes.len() - rest.len()).unwrap_or(u64::MAX)).saturating_add(needed),
)
.unwrap_or(usize::MAX),
found: bytes.len(),
};
let (count, rest) = read_u32(rest).ok_or_else(|| short(rest, 4))?;
if let Some(budget) = budget
&& u64::from(count) > budget.max_candidates() as u64
{
return Err(SealDecodeError::TooManyLedgerDots {
count: u64::from(count),
budget: budget.max_candidates() as u64,
});
}
let backable = rest.len().saturating_sub(HAVE_SET_WIRE_HEADER_LEN) / WIRE_DOT_LEN;
if count as usize > backable {
let floor = u64::from(count)
.saturating_mul(WIRE_DOT_LEN as u64)
.saturating_add(HAVE_SET_WIRE_HEADER_LEN as u64);
return Err(short(rest, floor));
}
let mut protocol = BTreeSet::new();
let mut previous: Option<Dot> = None;
let mut rest = rest;
for _ in 0..count {
let (raw, tail) = read_dot(rest).ok_or_else(|| short(rest, WIRE_DOT_LEN as u64))?;
let Ok(dot) = Dot::from_parts(raw.0, raw.1) else {
return Err(SealDecodeError::ZeroLedgerDot { station: raw.0 });
};
if let Some(previous) = previous
&& dot <= previous
{
return Err(SealDecodeError::NonAscendingLedger {
previous,
found: dot,
});
}
previous = Some(dot);
let _ = protocol.insert(dot);
rest = tail;
}
Ok((protocol, rest))
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct SealRecord {
declaration: EpochAddress,
candidates: BTreeSet<EpochAddress>,
protocol: BTreeSet<Dot>,
sealed_join: VersionVector,
admission: Option<Admission>,
}
impl SealRecord {
#[must_use]
pub fn from_sealed(sealed: &SealedEpoch) -> Self {
Self {
declaration: sealed.declaration,
candidates: sealed.candidates.clone(),
protocol: sealed.protocol.clone(),
sealed_join: sealed.sealed_join.clone(),
admission: sealed.admission.clone(),
}
}
#[must_use]
pub const fn admission(&self) -> Option<&Admission> {
self.admission.as_ref()
}
#[must_use]
pub const fn declaration(&self) -> EpochAddress {
self.declaration
}
#[must_use]
pub const fn sealed_join(&self) -> &VersionVector {
&self.sealed_join
}
pub fn candidates(&self) -> impl Iterator<Item = EpochAddress> + '_ {
self.candidates.iter().copied()
}
pub fn declaration_dots(&self) -> impl Iterator<Item = Dot> + '_ {
self.protocol.iter().copied()
}
#[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() {
SEAL_WIRE_V2
} else {
SEAL_WIRE_V1
});
encode_address_into(&mut out, self.declaration);
let candidate_count = u32::try_from(self.candidates.len()).unwrap_or(u32::MAX);
out.extend_from_slice(&candidate_count.to_be_bytes());
for candidate in &self.candidates {
encode_address_into(&mut out, *candidate);
}
let ledger_count = u32::try_from(self.protocol.len()).unwrap_or(u32::MAX);
out.extend_from_slice(&ledger_count.to_be_bytes());
for &dot in &self.protocol {
encode_dot_into(&mut out, dot);
}
encode_cut_into(&mut out, &self.sealed_join);
if let Some(admission) = &self.admission {
encode_address_into(&mut out, admission.base());
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());
}
}
out
}
#[must_use]
pub fn encoded_len(&self) -> usize {
1 + WIRE_ADDRESS_LEN
+ 4
+ self.candidates.len() * WIRE_ADDRESS_LEN
+ 4
+ self.protocol.len() * WIRE_DOT_LEN
+ cut_embedding_len(&self.sealed_join)
+ self.admission.as_ref().map_or(0, |admission| {
WIRE_ADDRESS_LEN + 4 + admission.joiners().count() * 4
})
}
pub fn from_bytes(bytes: &[u8]) -> Result<Self, SealDecodeError> {
let (record, tail) = Self::from_prefix(bytes)?;
Self::reject_tail(bytes, tail)?;
Ok(record)
}
pub fn from_bytes_with_budget(
bytes: &[u8],
budget: SealDecodeBudget,
) -> Result<Self, SealDecodeError> {
let (record, tail) = Self::from_prefix_with_budget(bytes, Some(budget))?;
Self::reject_tail(bytes, tail)?;
Ok(record)
}
pub fn from_prefix(bytes: &[u8]) -> Result<(Self, &[u8]), SealDecodeError> {
Self::from_prefix_with_budget(bytes, None)
}
pub(super) fn from_prefix_with_budget(
bytes: &[u8],
budget: Option<SealDecodeBudget>,
) -> Result<(Self, &[u8]), SealDecodeError> {
if bytes.len() < SEAL_WIRE_MIN_LEN {
return Err(SealDecodeError::UnexpectedLength {
expected: SEAL_WIRE_MIN_LEN,
found: bytes.len(),
});
}
let version = bytes[0];
if version != SEAL_WIRE_V1 && version != SEAL_WIRE_V2 {
return Err(SealDecodeError::UnknownVersion(version));
}
let ((generation, dot), rest) =
read_address_parts(&bytes[1..]).ok_or(SealDecodeError::UnexpectedLength {
expected: SEAL_WIRE_MIN_LEN,
found: bytes.len(),
})?;
let declaration = EpochAddress::try_from_parts(generation, dot)?;
let (candidates, rest) = read_candidate_set(bytes, rest, budget)?;
let (protocol, rest) = read_ledger_set(bytes, rest, budget)?;
let (sealed_join, tail) = decode_cut(rest)?;
let (admission, tail) = if version == SEAL_WIRE_V2 {
let (admission, tail) = read_admission(bytes, tail, declaration, budget)?;
(Some(admission), tail)
} else {
(None, tail)
};
Ok((
Self {
declaration,
candidates,
protocol,
sealed_join,
admission,
},
tail,
))
}
const fn reject_tail(bytes: &[u8], tail: &[u8]) -> Result<(), SealDecodeError> {
if tail.is_empty() {
Ok(())
} else {
Err(SealDecodeError::UnexpectedLength {
expected: bytes.len() - tail.len(),
found: bytes.len(),
})
}
}
}