use crate::model::ModelFormat;
use std::borrow::Cow;
use std::io::Read;
use super::BoostedModel;
use super::native::write_container_into;
use super::sections::{Sections, Writer, format_error};
use crate::error::Result;
pub(crate) const WRITER: &str = concat!("hessboost ", env!("CARGO_PKG_VERSION"));
pub(super) const ZSTD_MAGIC: [u8; 4] = [0x28, 0xB5, 0x2F, 0xFD];
pub(crate) fn is_zstd_frame(bytes: &[u8]) -> bool {
bytes.starts_with(&ZSTD_MAGIC)
}
const ALWAYS_ALLOWED: u64 = 256 << 20;
pub(super) const MAX_EXPANSION: u64 = 1 << 12;
pub(crate) struct ContainerSpec {
pub(crate) magic: [u8; 4],
pub(crate) version: u8,
pub(crate) what: &'static str,
pub(crate) known: fn(&str) -> bool,
pub(crate) legacy: Option<([u8; 4], &'static str)>,
}
impl ContainerSpec {
pub(crate) fn frame(&self, w: Writer) -> Vec<u8> {
let mut out = Vec::new();
self.frame_into(w, &mut out);
out
}
pub(crate) fn frame_into(&self, w: Writer, out: &mut Vec<u8>) {
let start = out.len();
out.reserve(self.magic.len() + 1 + w.encoded_len() + 8);
out.extend_from_slice(&self.magic);
out.push(self.version);
w.finish(out);
let checksum = xxh64(&out[start..]);
out.extend_from_slice(&checksum.to_le_bytes());
}
pub(crate) fn seal(&self, w: Writer) -> Result<Vec<u8>> {
let container = self.frame(w);
let frame = zstd::bulk::compress(&container, zstd::DEFAULT_COMPRESSION_LEVEL)?;
Ok(
if expansion_accepted(frame.len() as u64, container.len() as u64) {
frame
} else {
container
},
)
}
pub(crate) fn read<T>(
&self,
bytes: &[u8],
f: impl FnOnce(&Sections) -> Result<T>,
) -> Result<T> {
let container = unpack(bytes)?;
if let Some((magic, refusal)) = self.legacy
&& container.starts_with(&magic)
{
return Err(format_error(refusal));
}
let table = section_table(&container, self.magic, self.version, self.what)?;
let (s, rest) = Sections::parse(table, self.known)?;
if !rest.is_empty() {
return Err(format_error(format!(
"{} unexpected bytes after the last section",
rest.len()
)));
}
f(&s)
}
}
pub(crate) fn write_models(models: &[BoostedModel]) -> Result<(Vec<u64>, Vec<u8>)> {
let mut lengths = Vec::with_capacity(models.len());
let mut blob = Vec::new();
for m in models {
let start = blob.len();
write_container_into(m, &mut blob)?;
lengths.push((blob.len() - start) as u64);
}
Ok((lengths, blob))
}
pub(crate) fn read_models(
lengths: &[u64],
mut blob: &[u8],
lengths_name: &str,
blob_name: &str,
) -> Result<Vec<BoostedModel>> {
let mut models = Vec::with_capacity(lengths.len().min(blob.len()));
for &len in lengths {
let len = usize::try_from(len)
.ok()
.filter(|&len| len <= blob.len())
.ok_or_else(|| format_error(format!("section `{lengths_name}` is out of range")))?;
let (model, rest) = blob.split_at(len);
models.push(BoostedModel::decode(model, ModelFormat::Binary)?);
blob = rest;
}
if !blob.is_empty() {
return Err(format_error(format!(
"section `{blob_name}` has bytes past its models"
)));
}
Ok(models)
}
fn unpack(bytes: &[u8]) -> Result<Cow<'_, [u8]>> {
Ok(if bytes.starts_with(&ZSTD_MAGIC) {
Cow::Owned(decompress(bytes)?)
} else {
Cow::Borrowed(bytes)
})
}
fn section_table<'a>(
container: &'a [u8],
magic: [u8; 4],
version: u8,
what: &str,
) -> Result<&'a [u8]> {
let Some(body) = container.strip_prefix(&magic) else {
return Err(format_error(format!("invalid {what} header")));
};
let Some(&stored) = body.first() else {
return Err(format_error(format!("truncated {what}")));
};
if stored != version {
return Err(format_error(format!("unsupported {what} version {stored}")));
}
let Some(split) = container
.len()
.checked_sub(8)
.filter(|&at| at > magic.len())
else {
return Err(format_error(format!("truncated {what}")));
};
let (checked, checksum) = container.split_at(split);
if xxh64(checked).to_le_bytes() != checksum {
return Err(format_error(format!("{what} checksum mismatch")));
}
Ok(&checked[magic.len() + 1..])
}
fn expansion_limit(compressed: u64) -> u64 {
compressed.saturating_mul(MAX_EXPANSION).max(ALWAYS_ALLOWED)
}
fn expansion_accepted(compressed: u64, decompressed: u64) -> bool {
decompressed <= expansion_limit(compressed)
}
pub(super) fn decompress(bytes: &[u8]) -> Result<Vec<u8>> {
let limit = expansion_limit(bytes.len() as u64);
let decoder = zstd::stream::read::Decoder::with_buffer(bytes)
.map_err(|e| format_error(format!("zstd: {e}")))?;
let mut out = Vec::new();
decoder
.take(limit.saturating_add(1))
.read_to_end(&mut out)
.map_err(|e| format_error(format!("zstd: {e}")))?;
if out.len() as u64 > limit {
return Err(format_error("zstd frame expands too far"));
}
Ok(out)
}
pub(super) fn xxh64(data: &[u8]) -> u64 {
const P1: u64 = 0x9E37_79B1_85EB_CA87;
const P2: u64 = 0xC2B2_AE3D_27D4_EB4F;
const P3: u64 = 0x1656_67B1_9E37_79F9;
const P4: u64 = 0x85EB_CA77_C2B2_AE63;
const P5: u64 = 0x27D4_EB2F_1656_67C5;
let round = |acc: u64, lane: u64| {
acc.wrapping_add(lane.wrapping_mul(P2))
.rotate_left(31)
.wrapping_mul(P1)
};
let merge = |acc: u64, v: u64| (acc ^ round(0, v)).wrapping_mul(P1).wrapping_add(P4);
let (stripes, tail) = data.as_chunks::<32>();
let mut h = if stripes.is_empty() {
P5
} else {
let mut v = [P1.wrapping_add(P2), P2, 0, P1.wrapping_neg()];
for stripe in stripes {
for (lane, acc) in stripe.as_chunks::<8>().0.iter().zip(&mut v) {
*acc = round(*acc, u64::from_le_bytes(*lane));
}
}
let h = v[0]
.rotate_left(1)
.wrapping_add(v[1].rotate_left(7))
.wrapping_add(v[2].rotate_left(12))
.wrapping_add(v[3].rotate_left(18));
v.into_iter().fold(h, merge)
};
h = h.wrapping_add(data.len() as u64);
let (words, mut rest) = tail.as_chunks::<8>();
for word in words {
h = (h ^ round(0, u64::from_le_bytes(*word)))
.rotate_left(27)
.wrapping_mul(P1)
.wrapping_add(P4);
}
if let Some((half, after)) = rest.split_first_chunk::<4>() {
let half = u32::from_le_bytes(*half);
h = (h ^ u64::from(half).wrapping_mul(P1))
.rotate_left(23)
.wrapping_mul(P2)
.wrapping_add(P3);
rest = after;
}
for &byte in rest {
h = (h ^ u64::from(byte).wrapping_mul(P5))
.rotate_left(11)
.wrapping_mul(P1);
}
h ^= h >> 33;
h = h.wrapping_mul(P2);
h ^= h >> 29;
h = h.wrapping_mul(P3);
h ^ (h >> 32)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn xxh64_matches_the_reference() {
let bytes: Vec<u8> = (0..100).collect();
for (data, expected) in [
(&b""[..], 0xef46_db37_51d8_e999),
(b"a", 0xd24e_c4f1_a98c_6e5b),
(b"abc", 0x44bc_2cf5_ad77_0999),
(b"0123456789abcdef", 0x5c5b_90c3_4e37_6d0b),
(&bytes[..31], 0xc346_d2b5_9b4d_8ee1),
(&bytes[..32], 0xcbf5_9c51_16ff_32b4),
(&bytes[..], 0x6ac1_e580_3216_6597),
] {
assert_eq!(xxh64(data), expected, "{} bytes", data.len());
}
}
#[test]
fn frame_expansion_policy_boundaries() {
assert!(expansion_accepted(1, ALWAYS_ALLOWED));
assert!(!expansion_accepted(1, ALWAYS_ALLOWED + 1));
let frame = ALWAYS_ALLOWED / MAX_EXPANSION + 1;
assert!(expansion_accepted(frame, frame * MAX_EXPANSION));
assert!(!expansion_accepted(frame, frame * MAX_EXPANSION + 1));
assert!(expansion_accepted(u64::MAX, u64::MAX));
}
}