use super::error::{BinaryError, BinaryResult};
use std::fmt;
use std::io::{self, Read, Write};
use std::str::FromStr;
pub const BINARY_VERSION_1: u8 = 0x10;
const FLAG_EXTERNAL_REFERENCES: u8 = 0b0001;
const FLAG_EXPLICIT_LAYOUT: u8 = 0b0010;
const WIDTH_SHIFT: u8 = 2;
const WIDTH_BITS: u8 = 0b1100;
const SHAPE_WIDTH_BITS: u64 = 0b0011;
const SHAPE_HAS_GAP: u64 = 0b0100;
const SHAPE_VARIABLE_ARITY: u64 = 0b1000;
const SHAPE_MIN_ARITY_SHIFT: u32 = 4;
pub const NULL: u64 = 0;
pub const ONE: u64 = 1;
pub const NUMBER: u64 = 2;
pub const STRING: u64 = 3;
pub const LIST: u64 = 4;
pub const IDENTIFIED: u64 = 5;
pub const FIRST_LINK_ADDRESS: u64 = 6;
const COMPACT_GAP: u64 = FIRST_LINK_ADDRESS - 1;
pub const WIDTHS: [u8; 4] = [1, 2, 4, 8];
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub enum Reference {
Internal(u64),
External(u64),
}
impl Reference {
pub const NULL: Reference = Reference::Internal(NULL);
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct ArityRange {
pub min: u64,
pub max: Option<u64>,
}
impl ArityRange {
pub const DOUBLETS: ArityRange = ArityRange::exactly(2);
pub const fn exactly(arity: u64) -> Self {
Self {
min: arity,
max: Some(arity),
}
}
pub const fn at_least(min: u64) -> Self {
Self { min, max: None }
}
pub const fn between(min: u64, max: u64) -> Self {
Self {
min,
max: Some(max),
}
}
pub fn contains(&self, length: u64) -> bool {
length >= self.min && self.max.is_none_or(|max| length <= max)
}
pub fn is_fixed(&self) -> bool {
self.max == Some(self.min)
}
pub(crate) fn validate(&self) -> Result<(), String> {
if self.min == 0 {
return Err("arity must be at least 1".into());
}
if self.min > u64::MAX >> SHAPE_MIN_ARITY_SHIFT {
return Err(format!("arity {} is too large", self.min));
}
if self.max.is_some_and(|max| max < self.min) {
return Err(format!("arity range {self} is empty"));
}
Ok(())
}
fn extra(&self) -> u64 {
self.max.map_or(0, |max| max - self.min)
}
fn from_shape(min: u64, extra: Option<u64>) -> BinaryResult<Self> {
let range = match extra {
None => Self::exactly(min),
Some(0) => Self::at_least(min),
Some(extra) => Self::between(
min,
min.checked_add(extra)
.ok_or_else(|| BinaryError::malformed("arity range overflows 64 bits"))?,
),
};
range.validate().map_err(BinaryError::malformed)?;
Ok(range)
}
}
impl Default for ArityRange {
fn default() -> Self {
Self::DOUBLETS
}
}
impl fmt::Display for ArityRange {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self.max {
Some(max) if max == self.min => write!(formatter, "{max}"),
Some(max) => write!(formatter, "{}..{max}", self.min),
None => write!(formatter, "{}..", self.min),
}
}
}
impl FromStr for ArityRange {
type Err = String;
fn from_str(text: &str) -> Result<Self, Self::Err> {
let number = |part: &str| {
part.parse::<u64>()
.map_err(|_| format!("invalid arity '{text}': expected n, min..max or min.."))
};
let range = match text.split_once("..") {
None => Self::exactly(number(text)?),
Some((min, "")) => Self::at_least(number(min)?),
Some((min, max)) => Self::between(number(min)?, number(max)?),
};
range.validate()?;
Ok(range)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct DecodeLimits {
pub max_links: u64,
pub max_references: u64,
pub max_nodes: usize,
pub max_string_bytes: usize,
pub max_depth: usize,
}
impl Default for DecodeLimits {
fn default() -> Self {
Self {
max_links: 1 << 22,
max_references: 1 << 24,
max_nodes: 1 << 22,
max_string_bytes: 64 << 20,
max_depth: 64,
}
}
}
impl DecodeLimits {
pub fn unlimited() -> Self {
Self {
max_links: u64::MAX,
max_references: u64::MAX,
max_nodes: usize::MAX,
max_string_bytes: usize::MAX,
max_depth: usize::MAX,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Section {
pub gap: u64,
pub arity: ArityRange,
pub width: u8,
pub links: Vec<Vec<Reference>>,
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct LinksPacket {
pub external_references: bool,
pub sections: Vec<Section>,
}
pub fn internal_capacity(width: u8, external_references: bool) -> u64 {
let bits = u32::from(width) * 8 - u32::from(external_references);
if bits >= 64 {
u64::MAX
} else {
(1u64 << bits) - 1
}
}
pub fn external_capacity(width: u8) -> u64 {
(1u64 << (u32::from(width) * 8 - 1)) - 1
}
pub fn address_tier(address: u64, external_references: bool) -> BinaryResult<u8> {
WIDTHS
.into_iter()
.find(|&width| internal_capacity(width, external_references) >= address)
.ok_or_else(|| {
BinaryError::Unencodable(format!("address {address} exceeds the internal range"))
})
}
fn width_mask(width: u8) -> u64 {
internal_capacity(width, false)
}
pub fn encode_external(value: u64, width: u8) -> u64 {
if value == 0 {
1u64 << (u32::from(width) * 8 - 1)
} else {
value.wrapping_neg() & width_mask(width)
}
}
pub fn decode_external(raw: u64, width: u8) -> Option<u64> {
let external_zero = 1u64 << (u32::from(width) * 8 - 1);
if raw == external_zero {
Some(0)
} else if raw > external_zero {
Some(raw.wrapping_neg() & width_mask(width))
} else {
None
}
}
fn width_code(width: u8) -> BinaryResult<u8> {
WIDTHS
.iter()
.position(|&candidate| candidate == width)
.map(|code| code as u8)
.ok_or_else(|| BinaryError::Unencodable(format!("invalid width {width}")))
}
fn width_from_code(code: u64) -> u8 {
WIDTHS[(code & 0b11) as usize]
}
pub fn reference_width(reference: Reference, external_references: bool) -> BinaryResult<u8> {
match reference {
Reference::Internal(address) => address_tier(address, external_references),
Reference::External(value) => {
if !external_references {
return Err(BinaryError::Unencodable(
"external reference in a packet without external references".into(),
));
}
WIDTHS
.into_iter()
.find(|&width| external_capacity(width) >= value)
.ok_or_else(|| {
BinaryError::Unencodable(format!("external value {value} exceeds 63 bits"))
})
}
}
}
impl Section {
fn write_header(&self, out: &mut Vec<u8>) -> BinaryResult<()> {
let mut shape =
u64::from(width_code(self.width)?) | (self.arity.min << SHAPE_MIN_ARITY_SHIFT);
if self.gap != 0 {
shape |= SHAPE_HAS_GAP;
}
if !self.arity.is_fixed() {
shape |= SHAPE_VARIABLE_ARITY;
}
write_leb128(out, shape);
if self.gap != 0 {
write_leb128(out, self.gap);
}
if !self.arity.is_fixed() {
write_leb128(out, self.arity.extra());
}
write_leb128(out, self.links.len() as u64);
Ok(())
}
fn read_header(reader: &mut dyn Read) -> BinaryResult<(Self, u64)> {
let shape = read_leb128(reader)?;
let width = width_from_code(shape & SHAPE_WIDTH_BITS);
let gap = if shape & SHAPE_HAS_GAP != 0 {
read_leb128(reader)?
} else {
0
};
let extra = if shape & SHAPE_VARIABLE_ARITY != 0 {
Some(read_leb128(reader)?)
} else {
None
};
let arity = ArityRange::from_shape(shape >> SHAPE_MIN_ARITY_SHIFT, extra)?;
let count = read_leb128(reader)?;
let section = Section {
gap,
arity,
width,
links: Vec::new(),
};
Ok((section, count))
}
fn validate(&self) -> BinaryResult<()> {
width_code(self.width)?;
self.arity.validate().map_err(BinaryError::Unencodable)?;
for link in &self.links {
if !self.arity.contains(link.len() as u64) {
return Err(BinaryError::Unencodable(format!(
"a link of {} references in a section of arity {}",
link.len(),
self.arity
)));
}
}
Ok(())
}
}
impl LinksPacket {
pub fn new(external_references: bool) -> Self {
Self {
external_references,
sections: Vec::new(),
}
}
pub fn pack(
external_references: bool,
links: &[(u64, Vec<Reference>)],
packed_widths: bool,
) -> BinaryResult<Self> {
let planner = SectionPlanner::new(external_references, links)?;
let uniform_layout = planner.plan(false);
let uniform = Self::lay_out(external_references, links, &uniform_layout);
if !packed_widths {
return Ok(uniform);
}
let packed_layout = planner.plan(true);
if packed_layout == uniform_layout {
return Ok(uniform);
}
let packed = Self::lay_out(external_references, links, &packed_layout);
Ok(if packed.to_bytes()?.len() < uniform.to_bytes()?.len() {
packed
} else {
uniform
})
}
fn lay_out(
external_references: bool,
links: &[(u64, Vec<Reference>)],
layout: &[(usize, u8)],
) -> Self {
let mut sections = Vec::with_capacity(layout.len());
let mut next_address = 1u64;
let mut remaining = links;
for &(count, width) in layout {
let (members, rest) = remaining.split_at(count);
remaining = rest;
let start = members[0].0;
let (shortest, longest) = members
.iter()
.map(|(_, link)| link.len() as u64)
.fold((u64::MAX, 0), |(shortest, longest), length| {
(shortest.min(length), longest.max(length))
});
sections.push(Section {
gap: start - next_address,
arity: ArityRange::between(shortest, longest),
width,
links: members.iter().map(|(_, link)| link.clone()).collect(),
});
next_address = start + count as u64;
}
Self {
external_references,
sections,
}
}
pub fn links(&self) -> impl Iterator<Item = (u64, &[Reference])> + '_ {
let mut next_address = 1u64;
self.sections.iter().flat_map(move |section| {
let start = next_address.saturating_add(section.gap);
next_address = start.saturating_add(section.links.len() as u64);
section
.links
.iter()
.enumerate()
.map(move |(index, link)| (start + index as u64, link.as_slice()))
})
}
pub fn link_count(&self) -> u64 {
self.sections
.iter()
.map(|section| section.links.len() as u64)
.sum()
}
pub(crate) fn validate(&self) -> BinaryResult<()> {
let mut next_address = 1u64;
for section in &self.sections {
section.validate()?;
next_address = next_address
.checked_add(section.gap)
.and_then(|start| start.checked_add(section.links.len() as u64))
.ok_or_else(|| BinaryError::Unencodable("addresses overflow 64 bits".into()))?;
for link in §ion.links {
for &reference in link {
self.raw_reference(reference, section.width)?;
}
}
}
Ok(())
}
fn is_compact(&self) -> bool {
match self.sections.as_slice() {
[] => true,
[only] => {
only.gap == COMPACT_GAP
&& only.arity == ArityRange::DOUBLETS
&& !only.links.is_empty()
}
_ => false,
}
}
pub fn to_bytes(&self) -> BinaryResult<Vec<u8>> {
let mut bytes = Vec::new();
self.write_to(&mut bytes)?;
Ok(bytes)
}
pub fn write_to(&self, writer: &mut dyn Write) -> BinaryResult<()> {
let mut header = BINARY_VERSION_1;
if self.external_references {
header |= FLAG_EXTERNAL_REFERENCES;
}
let mut out = Vec::new();
self.validate()?;
if self.is_compact() {
let width = self.sections.first().map_or(1, |section| section.width);
header |= width_code(width)? << WIDTH_SHIFT;
out.push(header);
write_leb128(&mut out, self.link_count());
} else {
out.push(header | FLAG_EXPLICIT_LAYOUT);
write_leb128(&mut out, self.sections.len() as u64);
for section in &self.sections {
section.write_header(&mut out)?;
}
}
for section in &self.sections {
for link in §ion.links {
if !section.arity.is_fixed() {
write_leb128(&mut out, link.len() as u64 - section.arity.min);
}
for &reference in link {
write_raw(
&mut out,
self.raw_reference(reference, section.width)?,
section.width,
);
}
}
}
writer.write_all(&out)?;
Ok(())
}
fn raw_reference(&self, reference: Reference, width: u8) -> BinaryResult<u64> {
if reference_width(reference, self.external_references)? > width {
return Err(BinaryError::Unencodable(format!(
"{reference:?} does not fit {width} byte(s)"
)));
}
Ok(match reference {
Reference::Internal(address) => address,
Reference::External(value) => encode_external(value, width),
})
}
pub fn from_bytes(bytes: &[u8], limits: &DecodeLimits) -> BinaryResult<Self> {
let mut cursor = bytes;
let packet = Self::read_from(&mut cursor, limits)?
.ok_or_else(|| BinaryError::malformed("empty input"))?;
if !cursor.is_empty() {
return Err(BinaryError::malformed("trailing bytes after the packet"));
}
Ok(packet)
}
pub fn read_from(reader: &mut dyn Read, limits: &DecodeLimits) -> BinaryResult<Option<Self>> {
let Some(header) = read_byte_or_eof(reader)? else {
return Ok(None);
};
if header & 0xF0 != BINARY_VERSION_1 {
return Err(BinaryError::malformed(format!(
"unsupported binary header byte 0x{header:02X}"
)));
}
let mut packet = LinksPacket::new(header & FLAG_EXTERNAL_REFERENCES != 0);
let mut counts = Vec::new();
if header & FLAG_EXPLICIT_LAYOUT == 0 {
let width = width_from_code(u64::from((header & WIDTH_BITS) >> WIDTH_SHIFT));
let count = read_leb128(reader)?;
if count > u64::MAX - FIRST_LINK_ADDRESS {
return Err(BinaryError::malformed("addresses overflow 64 bits"));
}
if count > 0 {
packet.sections.push(Section {
gap: COMPACT_GAP,
arity: ArityRange::DOUBLETS,
width,
links: Vec::new(),
});
counts.push(count);
}
} else {
if header & WIDTH_BITS != 0 {
return Err(BinaryError::malformed(
"the explicit layout keeps the header width bits clear",
));
}
let section_count = read_leb128(reader)?;
if section_count > limits.max_links {
return Err(Self::too_many_links(limits));
}
let mut next_address = 1u64;
for _ in 0..section_count {
let (section, count) = Section::read_header(reader)?;
next_address = next_address
.checked_add(section.gap)
.and_then(|start| start.checked_add(count))
.ok_or_else(|| BinaryError::malformed("addresses overflow 64 bits"))?;
packet.sections.push(section);
counts.push(count);
}
}
counts
.iter()
.try_fold(0u64, |total, &count| total.checked_add(count))
.filter(|&total| total <= limits.max_links)
.ok_or_else(|| Self::too_many_links(limits))?;
let mut references_left = limits.max_references;
let external_references = packet.external_references;
for (section, count) in packet.sections.iter_mut().zip(counts) {
section.links.reserve(count.min(4096) as usize);
for _ in 0..count {
let length = if section.arity.is_fixed() {
section.arity.min
} else {
let length = read_leb128(reader)?
.checked_add(section.arity.min)
.filter(|&length| section.arity.contains(length))
.ok_or_else(|| {
BinaryError::malformed(format!(
"link length outside the section arity {}",
section.arity
))
})?;
length
};
references_left = references_left.checked_sub(length).ok_or_else(|| {
BinaryError::LimitExceeded(format!(
"references exceed the limit of {}",
limits.max_references
))
})?;
let mut link = Vec::with_capacity(length.min(4096) as usize);
for _ in 0..length {
let raw = read_raw(reader, section.width)?;
link.push(
match decode_external(raw, section.width).filter(|_| external_references) {
Some(value) => Reference::External(value),
None => Reference::Internal(raw),
},
);
}
section.links.push(link);
}
}
Ok(Some(packet))
}
fn too_many_links(limits: &DecodeLimits) -> BinaryError {
BinaryError::LimitExceeded(format!(
"packet declares more than {} links",
limits.max_links
))
}
}
const SECTION_HEADER_ESTIMATE: u64 = 2;
const VARIABLE_ARITY_ESTIMATE: u64 = 1;
const UNREACHABLE: u64 = u64::MAX / 4;
const STATES: usize = WIDTHS.len() * 2;
struct SectionPlanner<'a> {
links: &'a [(u64, Vec<Reference>)],
needs: Vec<u8>,
}
impl<'a> SectionPlanner<'a> {
fn new(external_references: bool, links: &'a [(u64, Vec<Reference>)]) -> BinaryResult<Self> {
let mut needs = Vec::with_capacity(links.len());
let mut previous_address = 0u64;
for (address, link) in links {
if *address <= previous_address {
return Err(BinaryError::Unencodable(format!(
"link addresses must ascend from 1, got {address} after {previous_address}"
)));
}
if *address == u64::MAX {
return Err(BinaryError::Unencodable(
"addresses overflow 64 bits".into(),
));
}
if link.is_empty() {
return Err(BinaryError::Unencodable(format!(
"link {address} has no references"
)));
}
previous_address = *address;
let mut need = 1u8;
for &reference in link {
need = need.max(reference_width(reference, external_references)?);
}
needs.push(need);
}
Ok(Self { links, needs })
}
fn first_cheapest(costs: &[u64; STATES]) -> usize {
let mut best = 0;
for state in 1..STATES {
if costs[state] < costs[best] {
best = state;
}
}
best
}
fn plan(&self, packed_widths: bool) -> Vec<(usize, u8)> {
if self.links.is_empty() {
return Vec::new();
}
let widest = self.needs.iter().copied().max().unwrap_or(1);
let allowed_widths = WIDTHS.map(|width| packed_widths || width == widest);
let mut opens_section = vec![0u8; self.links.len()];
let mut previous_best = vec![0u8; self.links.len()];
let mut costs = [UNREACHABLE; STATES];
for (index, (address, link)) in self.links.iter().enumerate() {
let best = Self::first_cheapest(&costs);
previous_best[index] = best as u8;
let cheapest_before = if index == 0 { 0 } else { costs[best] };
let continues = index > 0 && self.links[index - 1].0 + 1 == *address;
let same_length = index > 0 && self.links[index - 1].1.len() == link.len();
let mut next = [UNREACHABLE; STATES];
for state in 0..STATES {
let width_index = state / 2;
let variable = state % 2 == 1;
let width = WIDTHS[width_index];
if !allowed_widths[width_index] || width < self.needs[index] {
continue;
}
let variable_estimate = if variable { VARIABLE_ARITY_ESTIMATE } else { 0 };
let body = link.len() as u64 * u64::from(width) + variable_estimate;
let opening_cost = cheapest_before + SECTION_HEADER_ESTIMATE + variable_estimate;
let continuing_cost = if continues && (variable || same_length) {
costs[state]
} else {
UNREACHABLE
};
if continuing_cost <= opening_cost {
next[state] = continuing_cost + body;
} else {
next[state] = opening_cost + body;
opens_section[index] |= 1 << state;
}
}
costs = next;
}
let mut sections = Vec::new();
let mut state = Self::first_cheapest(&costs);
let mut end = self.links.len();
for index in (0..self.links.len()).rev() {
if opens_section[index] & (1 << state) != 0 {
sections.push((end - index, WIDTHS[state / 2]));
end = index;
state = usize::from(previous_best[index]);
}
}
sections.reverse();
sections
}
}
fn write_raw(out: &mut Vec<u8>, value: u64, width: u8) {
out.extend_from_slice(&value.to_le_bytes()[..usize::from(width)]);
}
fn read_raw(reader: &mut dyn Read, width: u8) -> BinaryResult<u64> {
let mut bytes = [0u8; 8];
read_exact(reader, &mut bytes[..usize::from(width)])?;
Ok(u64::from_le_bytes(bytes))
}
pub fn write_leb128(out: &mut Vec<u8>, mut value: u64) {
loop {
let byte = (value & 0x7F) as u8;
value >>= 7;
if value == 0 {
out.push(byte);
return;
}
out.push(byte | 0x80);
}
}
pub fn read_leb128(reader: &mut dyn Read) -> BinaryResult<u64> {
let mut value = 0u64;
for shift in (0..64).step_by(7) {
let mut byte = [0u8; 1];
read_exact(reader, &mut byte)?;
let payload = u64::from(byte[0] & 0x7F);
if shift == 63 && payload > 1 {
return Err(BinaryError::malformed("LEB128 value overflows 64 bits"));
}
value |= payload << shift;
if byte[0] & 0x80 == 0 {
return Ok(value);
}
}
Err(BinaryError::malformed("LEB128 value overflows 64 bits"))
}
fn read_exact(reader: &mut dyn Read, buffer: &mut [u8]) -> BinaryResult<()> {
reader.read_exact(buffer).map_err(|error| {
if error.kind() == io::ErrorKind::UnexpectedEof {
BinaryError::malformed("unexpected end of packet")
} else {
BinaryError::Io(error)
}
})
}
fn read_byte_or_eof(reader: &mut dyn Read) -> BinaryResult<Option<u8>> {
let mut byte = [0u8; 1];
loop {
match reader.read(&mut byte) {
Ok(0) => return Ok(None),
Ok(_) => return Ok(Some(byte[0])),
Err(error) if error.kind() == io::ErrorKind::Interrupted => {}
Err(error) => return Err(error.into()),
}
}
}