use std::collections::HashMap;
use crate::error::{HessboostError, Result};
pub(crate) const REQUIRED: u8 = 1;
#[derive(Default)]
pub(crate) struct Writer {
sections: Vec<(&'static str, u8, Vec<u8>)>,
}
impl Writer {
pub(crate) fn raw(&mut self, name: &'static str, flags: u8, payload: &[u8]) {
self.push(name, flags, payload.to_vec());
}
pub(crate) fn raw_owned(&mut self, name: &'static str, flags: u8, payload: Vec<u8>) {
self.push(name, flags, payload);
}
pub(crate) fn str(&mut self, name: &'static str, value: &str) {
self.raw(name, REQUIRED, value.as_bytes());
}
pub(crate) fn u64(&mut self, name: &'static str, value: u64) {
self.raw(name, REQUIRED, &value.to_le_bytes());
}
pub(crate) fn f64(&mut self, name: &'static str, value: f64) {
self.raw(name, REQUIRED, &value.to_le_bytes());
}
pub(crate) fn array<const N: usize, T: Copy>(
&mut self,
name: &'static str,
values: impl IntoIterator<Item = T>,
to_le: fn(T) -> [u8; N],
) {
let values = values.into_iter();
let mut payload = Vec::with_capacity(values.size_hint().0.saturating_mul(N));
for v in values {
payload.extend_from_slice(&to_le(v));
}
self.push(name, REQUIRED, payload);
}
fn push(&mut self, name: &'static str, flags: u8, payload: Vec<u8>) {
debug_assert!(u8::try_from(name.len()).is_ok());
self.sections.push((name, flags, payload));
}
pub(crate) fn encoded_len(&self) -> usize {
4 + self
.sections
.iter()
.map(|(name, _, payload)| 1 + name.len() + 1 + 8 + payload.len())
.sum::<usize>()
}
pub(crate) fn finish(self, out: &mut Vec<u8>) {
out.reserve(self.encoded_len());
out.extend_from_slice(&(self.sections.len() as u32).to_le_bytes());
for (name, flags, payload) in &self.sections {
out.push(name.len() as u8);
out.extend_from_slice(name.as_bytes());
out.push(*flags);
out.extend_from_slice(&(payload.len() as u64).to_le_bytes());
}
for (_, _, payload) in &self.sections {
out.extend_from_slice(payload);
}
}
}
pub(crate) struct Sections<'a>(HashMap<&'a str, &'a [u8]>);
impl<'a> Sections<'a> {
pub(crate) fn parse(bytes: &'a [u8], known: impl Fn(&str) -> bool) -> Result<(Self, &'a [u8])> {
let mut r = Reader(bytes);
let count = u32::from_le_bytes(r.array()?);
let mut table = Vec::new();
for _ in 0..count {
let name_len = usize::from(r.take(1)?[0]);
let name = std::str::from_utf8(r.take(name_len)?)
.map_err(|_| format_error("section name is not UTF-8"))?;
let flags = r.take(1)?[0];
let len = usize::try_from(u64::from_le_bytes(r.array()?))
.map_err(|_| format_error("section is too large"))?;
table.push((name, flags, len));
}
let mut sections = HashMap::with_capacity(table.len());
for (name, flags, len) in table {
let payload = r.take(len)?;
if !known(name) {
if flags & REQUIRED != 0 {
return Err(format_error(format!(
"the data needs section `{name}`, which this version cannot read"
)));
}
continue;
}
if sections.insert(name, payload).is_some() {
return Err(format_error(format!("section `{name}` appears twice")));
}
}
Ok((Sections(sections), r.0))
}
pub(crate) fn has(&self, name: &str) -> bool {
self.0.contains_key(name)
}
pub(crate) fn bytes(&self, name: &str) -> Result<&'a [u8]> {
self.0
.get(name)
.copied()
.ok_or_else(|| format_error(format!("missing section `{name}`")))
}
pub(crate) fn str(&self, name: &str) -> Result<&'a str> {
std::str::from_utf8(self.bytes(name)?)
.map_err(|_| format_error(format!("section `{name}` is not UTF-8")))
}
pub(crate) fn u64(&self, name: &str) -> Result<u64> {
Ok(u64::from_le_bytes(self.scalar(name)?))
}
pub(crate) fn usize(&self, name: &str) -> Result<usize> {
usize::try_from(self.u64(name)?)
.map_err(|_| format_error(format!("section `{name}` is out of range")))
}
pub(crate) fn f64(&self, name: &str) -> Result<f64> {
Ok(f64::from_le_bytes(self.scalar(name)?))
}
fn scalar(&self, name: &str) -> Result<[u8; 8]> {
self.bytes(name)?
.try_into()
.map_err(|_| format_error(format!("section `{name}` is not one value")))
}
pub(crate) fn array<const N: usize, T>(
&self,
name: &str,
from_le: fn([u8; N]) -> T,
) -> Result<Vec<T>> {
let (values, rest) = self.bytes(name)?.as_chunks::<N>();
if !rest.is_empty() {
return Err(format_error(format!(
"section `{name}` is not a whole number of values"
)));
}
Ok(values.iter().map(|&v| from_le(v)).collect())
}
pub(crate) fn bytes_exact(&self, name: &str, len: Option<usize>) -> Result<&'a [u8]> {
let bytes = self.bytes(name)?;
if Some(bytes.len()) == len {
Ok(bytes)
} else {
Err(wrong_length(name))
}
}
pub(crate) fn array_exact<const N: usize, T>(
&self,
name: &str,
count: usize,
from_le: fn([u8; N]) -> T,
) -> Result<Vec<T>> {
let bytes = self.bytes_exact(name, count.checked_mul(N))?;
Ok(bytes
.as_chunks::<N>()
.0
.iter()
.map(|&v| from_le(v))
.collect())
}
}
struct Reader<'a>(&'a [u8]);
impl<'a> Reader<'a> {
fn take(&mut self, n: usize) -> Result<&'a [u8]> {
if n > self.0.len() {
return Err(format_error("truncated data"));
}
let (head, rest) = self.0.split_at(n);
self.0 = rest;
Ok(head)
}
fn array<const N: usize>(&mut self) -> Result<[u8; N]> {
let mut out = [0; N];
out.copy_from_slice(self.take(N)?);
Ok(out)
}
}
pub(crate) fn format_error(msg: impl Into<String>) -> HessboostError {
HessboostError::model_format(msg)
}
pub(crate) fn wrong_length(name: &str) -> HessboostError {
format_error(format!("section `{name}` has the wrong length"))
}
pub(crate) fn unknown_value(name: &str, value: &str) -> HessboostError {
format_error(format!("unknown `{name}` value `{value}`"))
}