use super::container::{PositionalReader, Section, SectionReader};
use super::{validate_offsets, IndexSource, Input};
use anyhow::{ensure, Context, Result};
use std::io::{BufReader, Read, Seek, SeekFrom};
use std::path::Path;
use std::sync::Mutex;
const MAGIC: &[u8; 8] = b"SNDSP002";
#[cfg(feature = "import")]
const DICTIONARY_BYTES: usize = 112 * 1024;
#[cfg(feature = "import")]
const SAMPLE_BYTES: usize = 100 * DICTIONARY_BYTES;
#[cfg(feature = "import")]
const MINIMUM_SAMPLES: usize = 1000;
#[cfg(feature = "import")]
const LEVEL: i32 = 12;
pub struct DisplayStore {
positional: Option<PositionalReader>,
input: Mutex<BufReader<SectionReader>>,
offsets: Vec<u32>,
frames: Frames,
start: u64,
}
struct Frames {
lengths: Vec<u16>,
dictionary: Vec<u8>,
decoders: Mutex<Vec<zstd::bulk::Decompressor<'static>>>,
}
impl Frames {
fn decode(&self, frame: &[u8], length: usize) -> Result<Vec<u8>> {
let idle = self
.decoders
.lock()
.map_err(|_| anyhow::anyhow!("Display decoders poisoned"))?
.pop();
let mut decoder = match idle {
Some(decoder) => decoder,
None => zstd::bulk::Decompressor::with_dictionary(&self.dictionary)?,
};
let mut bytes = Vec::with_capacity(length);
decoder
.decompress_to_buffer(frame, &mut bytes)
.context("Invalid display label")?;
ensure!(bytes.len() == length, "Display label length mismatch");
if let Ok(mut idle) = self.decoders.lock() {
idle.push(decoder);
}
Ok(bytes)
}
}
impl DisplayStore {
pub(super) fn verify_text(&mut self) -> Result<()> {
self.all_labels().map(drop)
}
fn all_labels(&mut self) -> Result<Vec<Option<String>>> {
let input = self.input.get_mut().expect("display reader poisoned");
input.seek(SeekFrom::Start(self.start))?;
let mut text = Vec::new();
input.read_to_end(&mut text)?;
(0..self.offsets.len() - 1)
.map(|i| {
let stored = &text[self.offsets[i] as usize..self.offsets[i + 1] as usize];
self.label(i, stored)
})
.collect()
}
fn label(&self, i: usize, stored: &[u8]) -> Result<Option<String>> {
if stored.is_empty() {
return Ok(None);
}
let bytes = self
.frames
.decode(stored, self.frames.lengths[i] as usize)?;
Ok(Some(
String::from_utf8(bytes).context("Invalid display UTF-8")?,
))
}
#[cfg(feature = "import")]
pub(crate) fn into_labels(mut self) -> Result<Vec<Option<String>>> {
self.all_labels()
}
#[cfg(feature = "import")]
pub fn write(path: &Path, labels: &[Option<String>]) -> Result<()> {
use super::{put_u32s, put_u64};
use std::io::{BufWriter, Write};
let present: Vec<&[u8]> = labels.iter().flatten().map(|s| s.as_bytes()).collect();
let dictionary = if present.len() < MINIMUM_SAMPLES {
Vec::new()
} else {
let total: usize = present.iter().map(|l| l.len()).sum();
let step = total.div_ceil(SAMPLE_BYTES).max(1);
let mut sample = Vec::new();
let mut sizes = Vec::new();
for label in present.iter().step_by(step) {
sample.extend_from_slice(label);
sizes.push(label.len());
}
zstd::dict::from_continuous(&sample, &sizes, DICTIONARY_BYTES)
.context("Could not train the display dictionary")?
};
let mut compressor = zstd::bulk::Compressor::with_dictionary(LEVEL, &dictionary)?;
compressor.include_checksum(false)?;
compressor.include_contentsize(false)?;
compressor.include_dictid(false)?;
let mut offsets = Vec::with_capacity(labels.len() + 1);
let mut lengths = Vec::with_capacity(labels.len());
let mut frames = Vec::new();
offsets.push(0u32);
for label in labels {
let bytes = label.as_deref().unwrap_or_default().as_bytes();
lengths.push(u16::try_from(bytes.len()).context("Display label too long")?);
if !bytes.is_empty() {
frames.extend_from_slice(&compressor.compress(bytes)?);
}
offsets
.push(u32::try_from(frames.len()).context("Display section exceeds u32 capacity")?);
}
let mut out = BufWriter::new(std::fs::File::create_new(path)?);
out.write_all(MAGIC)?;
put_u32s(&mut out, &offsets)?;
put_u64(&mut out, lengths.len() as u64)?;
for length in &lengths {
out.write_all(&length.to_le_bytes())?;
}
put_u64(&mut out, dictionary.len() as u64)?;
out.write_all(&dictionary)?;
out.write_all(&frames)?;
out.flush()?;
out.get_ref().sync_all()?;
Ok(())
}
pub fn open(directory: &Path) -> Result<Self> {
let (manifest, source) = IndexSource::open(directory)?;
Self::from_section(&source.section("display.bin")?, manifest.concept_count)
}
fn from_section(section: &Section, concept_count: usize) -> Result<Self> {
let mut input = Input::open(section, MAGIC)?;
let offsets = input.u32s()?;
let frames = {
let lengths = input.u16s()?;
ensure!(
lengths.len() == concept_count
&& lengths
.iter()
.zip(offsets.windows(2))
.all(|(&length, w)| (length == 0) == (w[0] == w[1])),
"Invalid display lengths"
);
let dictionary = input.bytes()?;
Frames {
lengths,
dictionary,
decoders: Mutex::new(Vec::new()),
}
};
validate_offsets(&offsets, concept_count, input.remaining as usize)?;
let start = input.reader.stream_position()?;
Ok(Self {
positional: section.positional()?,
input: Mutex::new(input.reader),
offsets,
frames,
start,
})
}
pub fn prefetch(&self) -> Result<()> {
match &self.positional {
Some(reader) => reader.prefetch(),
None => Ok(()),
}
}
pub fn label_bytes(&self, ordinal: u32) -> Option<u32> {
let i = ordinal as usize;
self.frames.lengths.get(i).map(|&n| u32::from(n))
}
pub fn get(&self, ordinal: u32) -> Result<Option<String>> {
let i = ordinal as usize;
ensure!(i + 1 < self.offsets.len(), "Display ordinal out of range");
let length = (self.offsets[i + 1] - self.offsets[i]) as usize;
if length == 0 {
return Ok(None);
}
let position = self.start + self.offsets[i] as u64;
let mut stored = vec![0; length];
if let Some(reader) = &self.positional {
reader.read_exact_at(position, &mut stored)?;
} else {
let mut input = self
.input
.lock()
.map_err(|_| anyhow::anyhow!("Display reader poisoned"))?;
input.seek(SeekFrom::Start(position))?;
input.read_exact(&mut stored)?;
}
self.label(i, &stored)
}
}
#[cfg(all(test, feature = "import"))]
mod tests {
use super::*;
fn open(path: &Path) -> Result<DisplayStore> {
let length = std::fs::metadata(path)?.len();
DisplayStore::from_section(&Section::for_test(path, length, String::new()), 2001)
}
#[test]
fn labels_round_trip_through_a_trained_dictionary() {
let directory = tempfile::tempdir().unwrap();
let path = directory.path().join("display.bin");
let labels: Vec<Option<String>> = (0..2001)
.map(|i| {
(i != 7).then(|| format!("Synthetic finding {i} of left é structure (disorder)"))
})
.collect();
DisplayStore::write(&path, &labels).unwrap();
let store = open(&path).unwrap();
assert!(!store.frames.dictionary.is_empty());
for (i, label) in labels.iter().enumerate() {
assert_eq!(&store.get(i as u32).unwrap(), label);
assert_eq!(
store.label_bytes(i as u32),
Some(label.as_ref().map_or(0, |l| l.len() as u32))
);
}
assert!(store.get(2001).is_err());
assert_eq!(open(&path).unwrap().into_labels().unwrap(), labels);
let mut bytes = std::fs::read(&path).unwrap();
let last = bytes.len() - 2;
bytes[last] ^= 0xFF;
std::fs::write(&path, &bytes).unwrap();
let damaged = open(&path).unwrap();
assert!((0..2001).any(|i| damaged.get(i).is_err()));
}
}