use limnifs_core::codec::zstd_dict::{
compress_with_dict, decompress_with_dict, train_dictionary as core_train,
train_dictionary_fastcover,
};
pub const DEFAULT_TARGET_SIZE: usize = 65_536;
pub const DEFAULT_MIN_SAMPLES: usize = 100;
#[derive(Clone, Debug)]
pub struct TrainedDictionary {
pub id: u8,
pub codec: u8,
pub content: Vec<u8>,
}
impl TrainedDictionary {
pub fn compress(&self, plaintext: &[u8]) -> Result<Vec<u8>, crate::WriteError> {
compress_with_dict(plaintext, &self.content).map_err(|e| {
crate::WriteError::Io(std::io::Error::other(format!(
"dict compress (id {}): {e}",
self.id
)))
})
}
pub fn decompress(
&self,
compressed: &[u8],
expected_len: u32,
) -> Result<Vec<u8>, crate::WriteError> {
decompress_with_dict(compressed, expected_len, &self.content).map_err(|e| {
crate::WriteError::Io(std::io::Error::other(format!(
"dict decompress (id {}): {e}",
self.id
)))
})
}
}
#[must_use]
pub fn train_zstd(id: u8, samples: &[&[u8]], target_size: usize) -> Option<TrainedDictionary> {
train_zstd_with_trainer(id, samples, target_size, TrainerKind::Frequency)
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum TrainerKind {
Frequency,
FastCover,
}
impl TrainerKind {
#[must_use]
pub fn from_config_str(s: &str) -> Self {
match s.to_ascii_lowercase().as_str() {
"fastcover" => Self::FastCover,
_ => Self::Frequency,
}
}
}
#[must_use]
pub fn train_zstd_with_trainer(
id: u8,
samples: &[&[u8]],
target_size: usize,
trainer: TrainerKind,
) -> Option<TrainedDictionary> {
if samples.is_empty() || target_size == 0 {
return None;
}
let content = match trainer {
TrainerKind::Frequency => core_train(samples, target_size),
TrainerKind::FastCover => train_dictionary_fastcover(samples, target_size),
};
if content.is_empty() {
return None;
}
Some(TrainedDictionary {
id,
codec: limnifs_core::codec::CODEC_ZSTD,
content,
})
}
pub fn allocate_ids<'a>(class_names: &'a [&'a str]) -> Result<Vec<(&'a str, u8)>, &'static str> {
if class_names.len() > 254 {
return Err("dictionary id space exhausted (max 254 classes)");
}
Ok(class_names
.iter()
.enumerate()
.map(|(i, name)| (*name, u8::try_from(i).expect("≤ 254")))
.collect())
}
#[cfg(test)]
mod tests {
use super::*;
fn synthetic_text_samples(n: usize) -> Vec<Vec<u8>> {
(0..n)
.map(|i| format!("function test_case_{i}() {{ return {i}; }}\n").into_bytes())
.collect()
}
#[test]
fn train_zstd_returns_dict_for_repetitive_samples() {
let samples_vec = synthetic_text_samples(50);
let samples: Vec<&[u8]> = samples_vec.iter().map(Vec::as_slice).collect();
let dict = train_zstd(0, &samples, 4096);
if let Some(d) = &dict {
assert!(!d.content.is_empty(), "trained dict content non-empty");
assert_eq!(d.id, 0);
assert_eq!(d.codec, limnifs_core::codec::CODEC_ZSTD);
}
}
#[test]
fn train_zstd_returns_none_for_empty_samples() {
assert!(train_zstd(0, &[], 4096).is_none());
}
#[test]
fn train_zstd_returns_none_for_zero_target_size() {
let samples_vec = synthetic_text_samples(10);
let samples: Vec<&[u8]> = samples_vec.iter().map(Vec::as_slice).collect();
assert!(train_zstd(0, &samples, 0).is_none());
}
#[test]
fn dict_round_trips_when_trained() {
let samples_vec = synthetic_text_samples(50);
let samples: Vec<&[u8]> = samples_vec.iter().map(Vec::as_slice).collect();
let Some(dict) = train_zstd(0, &samples, 4096) else {
return; };
let plaintext = b"function test_case_99() { return 99; }\n";
let compressed = dict.compress(plaintext).expect("compress");
let recovered = dict
.decompress(&compressed, plaintext.len() as u32)
.expect("decompress");
assert_eq!(recovered.as_slice(), &plaintext[..]);
}
#[test]
fn allocate_ids_assigns_sequential_ids() {
let names = vec!["text", "binary", "source"];
let allocated = allocate_ids(&names).expect("allocate");
assert_eq!(allocated.len(), 3);
assert_eq!(allocated[0], ("text", 0));
assert_eq!(allocated[1], ("binary", 1));
assert_eq!(allocated[2], ("source", 2));
}
#[test]
fn allocate_ids_rejects_more_than_254_classes() {
let names: Vec<&str> = (0..255).map(|_| "x").collect();
assert!(allocate_ids(&names).is_err());
}
}