use std::fs::{self, File};
use std::io::{self, Write};
use std::path::{Path, PathBuf};
use memmap2::Mmap;
use crate::seam::{PreloadTable, PriorLookup, load_byte_ngram};
const SLOT: usize = 16;
const MANIFEST: &str = "prior.manifest";
const FORMAT: &str = "trex-prior-cache 1";
struct MappedOrder {
order: usize,
map: Mmap,
cap: usize,
zero_row: Option<(u32, u32)>,
}
pub struct MappedPrior {
orders: Vec<MappedOrder>,
}
fn blob_hash(blob: &[u8]) -> u64 {
blob.iter().fold(0xcbf2_9ce4_8422_2325u64, |h, &b| (h ^ u64::from(b)).wrapping_mul(0x0100_0000_01b3))
}
fn capacity_for(rows: usize) -> usize {
((rows * 10) / 7).next_power_of_two().max(16)
}
fn table_path(dir: &Path, order: usize) -> PathBuf {
dir.join(format!("order{order}.tbl"))
}
fn bad(msg: String) -> io::Error {
io::Error::new(io::ErrorKind::InvalidData, msg)
}
fn field(words: &[&str], i: usize, what: &str, line: &str) -> io::Result<usize> {
match words.get(i) {
None => Err(bad(format!("manifest line {line:?} has no {what}"))),
Some(s) => match s.parse::<usize>() {
Ok(v) => Ok(v),
Err(e) => Err(bad(format!("manifest line {line:?}: {what} {s:?} does not parse: {e}"))),
},
}
}
fn narrow(v: usize, what: &str, line: &str) -> io::Result<u32> {
match u32::try_from(v) {
Ok(n) => Ok(n),
Err(e) => Err(bad(format!("manifest line {line:?}: {what} {v} does not fit a count: {e}"))),
}
}
impl MappedPrior {
pub fn build(dir: &Path, blob: &[u8]) -> io::Result<Self> {
fs::create_dir_all(dir)?;
let tables: Vec<PreloadTable> = load_byte_ngram(blob);
let mut manifest = format!("{FORMAT}\nblob-len {}\nblob-hash {:016x}\n", blob.len(), blob_hash(blob));
for (order, rows) in &tables {
let cap = capacity_for(rows.len());
let mask = cap - 1;
let mut table = vec![0u8; cap * SLOT];
let mut zero_row = None;
for (&key, &(a, b)) in rows {
if key == 0 {
zero_row = Some((a, b));
continue;
}
let mut i = (key as usize) & mask;
loop {
let off = i * SLOT;
let here = u64::from_le_bytes(table[off..off + 8].try_into().expect("eight bytes"));
if here == 0 {
table[off..off + 8].copy_from_slice(&key.to_le_bytes());
table[off + 8..off + 12].copy_from_slice(&a.to_le_bytes());
table[off + 12..off + 16].copy_from_slice(&b.to_le_bytes());
break;
}
assert!(here != key, "the decoded table holds order {order} key {key:#x} once");
i = (i + 1) & mask;
}
}
let path = table_path(dir, *order);
let mut f = File::create(&path)?;
f.write_all(&table)?;
f.sync_all()?;
manifest.push_str(&format!("order {order} cap {cap} rows {}", rows.len()));
if let Some((a, b)) = zero_row {
manifest.push_str(&format!(" zero {a} {b}"));
}
manifest.push('\n');
}
let mut f = File::create(dir.join(MANIFEST))?;
f.write_all(manifest.as_bytes())?;
f.sync_all()?;
match Self::open(dir, blob)? {
Some(prior) => Ok(prior),
None => Err(bad(format!("the cache just written to {} does not describe this blob", dir.display()))),
}
}
#[allow(unsafe_code)]
pub fn open(dir: &Path, blob: &[u8]) -> io::Result<Option<Self>> {
let manifest_path = dir.join(MANIFEST);
let text = match fs::read_to_string(&manifest_path) {
Ok(t) => t,
Err(e) if e.kind() == io::ErrorKind::NotFound => return Ok(None),
Err(e) => return Err(e),
};
let mut lines = text.lines();
let mut head = |what: &str| -> io::Result<&str> {
match lines.next() {
Some(l) => Ok(l),
None => Err(bad(format!("{} ends before its {what}", manifest_path.display()))),
}
};
if head("format line")? != FORMAT {
return Err(bad(format!("{} is not a {FORMAT} manifest", manifest_path.display())));
}
let len_line = head("blob length")?;
let hash_line = head("blob hash")?;
if !len_line.starts_with("blob-len ") || !hash_line.starts_with("blob-hash ") {
return Err(bad(format!("{} lacks the blob length and hash", manifest_path.display())));
}
if len_line != format!("blob-len {}", blob.len()) || hash_line != format!("blob-hash {:016x}", blob_hash(blob)) {
return Ok(None);
}
let mut orders = Vec::new();
for line in lines.filter(|l| !l.trim().is_empty()) {
let w: Vec<&str> = line.split_whitespace().collect();
if w.first() != Some(&"order") || w.get(2) != Some(&"cap") || w.get(4) != Some(&"rows") {
return Err(bad(format!("{}: unreadable line {line:?}", manifest_path.display())));
}
let order = field(&w, 1, "order", line)?;
let cap = field(&w, 3, "cap", line)?;
if !cap.is_power_of_two() {
return Err(bad(format!("{}: order {order} cap {cap} is not a power of two", manifest_path.display())));
}
let zero_row = match w.get(6) {
Some(&"zero") => Some((
narrow(field(&w, 7, "zero n0", line)?, "zero n0", line)?,
narrow(field(&w, 8, "zero n1", line)?, "zero n1", line)?,
)),
Some(other) => return Err(bad(format!("{}: unexpected {other:?} in {line:?}", manifest_path.display()))),
None => None,
};
let path = table_path(dir, order);
let f = File::open(&path)?;
let len = f.metadata()?.len();
if len != (cap * SLOT) as u64 {
return Err(bad(format!("{} is {len} bytes, the manifest says {}", path.display(), cap * SLOT)));
}
let map = unsafe { Mmap::map(&f)? };
orders.push(MappedOrder { order, map, cap, zero_row });
}
Ok(Some(Self { orders }))
}
pub fn open_or_build(dir: &Path, blob: &[u8]) -> io::Result<(Self, bool)> {
match Self::open(dir, blob)? {
Some(prior) => Ok((prior, false)),
None => Ok((Self::build(dir, blob)?, true)),
}
}
#[must_use]
pub fn bytes_on_disk(&self) -> usize {
self.orders.iter().map(|o| o.cap * SLOT).sum()
}
#[must_use]
pub fn orders(&self) -> Vec<usize> {
self.orders.iter().map(|o| o.order).collect()
}
}
impl PriorLookup for MappedPrior {
fn has_order(&self, order: usize) -> bool {
self.orders.iter().any(|o| o.order == order)
}
fn get(&self, order: usize, key: u64) -> Option<(u32, u32)> {
let t = self.orders.iter().find(|o| o.order == order)?;
if key == 0 {
return t.zero_row;
}
let mask = t.cap - 1;
let mut i = (key as usize) & mask;
loop {
let off = i * SLOT;
let here = u64::from_le_bytes(t.map[off..off + 8].try_into().expect("eight bytes"));
if here == 0 {
return None;
}
if here == key {
let a = u32::from_le_bytes(t.map[off + 8..off + 12].try_into().expect("four bytes"));
let b = u32::from_le_bytes(t.map[off + 12..off + 16].try_into().expect("four bytes"));
return Some((a, b));
}
i = (i + 1) & mask;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::seam::{byte_ngram_train, logistic_mix_bits, logistic_mix_bits_with, serialize_byte_ngram};
const TEXT: &[u8] = b"It was the best of times, it was the worst of times, it was the age of wisdom, \
it was the age of foolishness, it was the epoch of belief, it was the epoch of incredulity, \
it was the season of Light, it was the season of Darkness, it was the spring of hope, \
it was the winter of despair. 1859, 1859, 1859 and again 1859.";
fn scratch(name: &str) -> PathBuf {
std::env::temp_dir().join(format!("trex/prior-cache-{name}-{}", std::process::id()))
}
#[test]
fn the_mapped_prior_answers_every_row_the_decoded_one_holds() {
let blob = serialize_byte_ngram(&byte_ngram_train(TEXT), 100_000);
let heap = load_byte_ngram(&blob);
let dir = scratch("rows");
let (mapped, built) = MappedPrior::open_or_build(&dir, &blob).expect("the cache builds");
assert!(built, "a fresh directory holds no cache");
let mut rows = 0usize;
for (order, table) in &heap {
assert!(mapped.has_order(*order));
for (&key, &counts) in table {
assert_eq!(mapped.get(*order, key), Some(counts), "order {order} key {key:#x}");
rows += 1;
}
assert_eq!(mapped.get(*order, 0x5eed_0000_dead_beef), None, "an absent key is absent");
}
assert!(rows > 100, "the sample trained {rows} rows, too few to test placement");
assert!(!mapped.has_order(7) && mapped.get(7, 1).is_none(), "order 7 is not a model order");
let (again, built_again) = MappedPrior::open_or_build(&dir, &blob).expect("the cache opens");
assert!(!built_again, "the second call opened the cache it found");
assert_eq!(again.orders(), mapped.orders());
assert_eq!(again.bytes_on_disk(), mapped.bytes_on_disk());
drop((mapped, again));
fs::remove_dir_all(&dir).expect("the scratch cache is removed");
}
#[test]
fn the_coder_codes_to_the_same_bits_against_either_form() {
let blob = serialize_byte_ngram(&byte_ngram_train(TEXT), 100_000);
let heap = load_byte_ngram(&blob);
let dir = scratch("bits");
let mapped = MappedPrior::build(&dir, &blob).expect("the cache builds");
let input = b"it was the age of hope and the season of 1859; the worst of belief was the epoch of wisdom.";
let from_heap = logistic_mix_bits(input, true, &[], Some(heap.as_slice()));
let from_mapped = logistic_mix_bits_with(input, true, &[], Some(&mapped));
let without = logistic_mix_bits(input, true, &[], None);
assert!(from_heap == from_mapped, "heap {from_heap} bits, mapped {from_mapped}");
assert!(from_heap != without, "the prior must change the bits, or the equality asserts nothing");
drop(mapped);
fs::remove_dir_all(&dir).expect("the scratch cache is removed");
}
#[test]
fn a_cache_from_another_blob_is_not_opened() {
let blob = serialize_byte_ngram(&byte_ngram_train(TEXT), 100_000);
let other = serialize_byte_ngram(&byte_ngram_train(&TEXT[..200]), 100_000);
let dir = scratch("stale");
let first = MappedPrior::build(&dir, &blob).expect("the cache builds");
drop(first);
assert!(MappedPrior::open(&dir, &other).expect("the manifest reads").is_none());
assert!(MappedPrior::open(&dir, &blob).expect("the manifest reads").is_some());
assert!(MappedPrior::open(&scratch("absent"), &blob).expect("no directory reads as no cache").is_none());
fs::remove_dir_all(&dir).expect("the scratch cache is removed");
}
}