limnifs_write/
dictionary.rs1use limnifs_core::codec::zstd_dict::{
33 compress_with_dict, decompress_with_dict, train_dictionary as core_train,
34 train_dictionary_fastcover,
35};
36
37pub const DEFAULT_TARGET_SIZE: usize = 65_536;
39
40pub const DEFAULT_MIN_SAMPLES: usize = 100;
44
45#[derive(Clone, Debug)]
47pub struct TrainedDictionary {
48 pub id: u8,
51 pub codec: u8,
53 pub content: Vec<u8>,
55}
56
57impl TrainedDictionary {
58 pub fn compress(&self, plaintext: &[u8]) -> Result<Vec<u8>, crate::WriteError> {
63 compress_with_dict(plaintext, &self.content).map_err(|e| {
64 crate::WriteError::Io(std::io::Error::other(format!(
65 "dict compress (id {}): {e}",
66 self.id
67 )))
68 })
69 }
70
71 pub fn decompress(
76 &self,
77 compressed: &[u8],
78 expected_len: u32,
79 ) -> Result<Vec<u8>, crate::WriteError> {
80 decompress_with_dict(compressed, expected_len, &self.content).map_err(|e| {
81 crate::WriteError::Io(std::io::Error::other(format!(
82 "dict decompress (id {}): {e}",
83 self.id
84 )))
85 })
86 }
87}
88
89#[must_use]
96pub fn train_zstd(id: u8, samples: &[&[u8]], target_size: usize) -> Option<TrainedDictionary> {
97 train_zstd_with_trainer(id, samples, target_size, TrainerKind::Frequency)
98}
99
100#[derive(Clone, Copy, Debug, Eq, PartialEq)]
102pub enum TrainerKind {
103 Frequency,
106 FastCover,
110}
111
112impl TrainerKind {
113 #[must_use]
116 pub fn from_config_str(s: &str) -> Self {
117 match s.to_ascii_lowercase().as_str() {
118 "fastcover" => Self::FastCover,
119 _ => Self::Frequency,
120 }
121 }
122}
123
124#[must_use]
127pub fn train_zstd_with_trainer(
128 id: u8,
129 samples: &[&[u8]],
130 target_size: usize,
131 trainer: TrainerKind,
132) -> Option<TrainedDictionary> {
133 if samples.is_empty() || target_size == 0 {
134 return None;
135 }
136 let content = match trainer {
137 TrainerKind::Frequency => core_train(samples, target_size),
138 TrainerKind::FastCover => train_dictionary_fastcover(samples, target_size),
139 };
140 if content.is_empty() {
141 return None;
142 }
143 Some(TrainedDictionary {
144 id,
145 codec: limnifs_core::codec::CODEC_ZSTD,
146 content,
147 })
148}
149
150pub fn allocate_ids<'a>(class_names: &'a [&'a str]) -> Result<Vec<(&'a str, u8)>, &'static str> {
156 if class_names.len() > 254 {
157 return Err("dictionary id space exhausted (max 254 classes)");
158 }
159 Ok(class_names
160 .iter()
161 .enumerate()
162 .map(|(i, name)| (*name, u8::try_from(i).expect("≤ 254")))
163 .collect())
164}
165
166#[cfg(test)]
167mod tests {
168 use super::*;
169
170 fn synthetic_text_samples(n: usize) -> Vec<Vec<u8>> {
171 (0..n)
174 .map(|i| format!("function test_case_{i}() {{ return {i}; }}\n").into_bytes())
175 .collect()
176 }
177
178 #[test]
179 fn train_zstd_returns_dict_for_repetitive_samples() {
180 let samples_vec = synthetic_text_samples(50);
181 let samples: Vec<&[u8]> = samples_vec.iter().map(Vec::as_slice).collect();
182 let dict = train_zstd(0, &samples, 4096);
183 if let Some(d) = &dict {
186 assert!(!d.content.is_empty(), "trained dict content non-empty");
187 assert_eq!(d.id, 0);
188 assert_eq!(d.codec, limnifs_core::codec::CODEC_ZSTD);
189 }
190 }
191
192 #[test]
193 fn train_zstd_returns_none_for_empty_samples() {
194 assert!(train_zstd(0, &[], 4096).is_none());
195 }
196
197 #[test]
198 fn train_zstd_returns_none_for_zero_target_size() {
199 let samples_vec = synthetic_text_samples(10);
200 let samples: Vec<&[u8]> = samples_vec.iter().map(Vec::as_slice).collect();
201 assert!(train_zstd(0, &samples, 0).is_none());
202 }
203
204 #[test]
205 fn dict_round_trips_when_trained() {
206 let samples_vec = synthetic_text_samples(50);
207 let samples: Vec<&[u8]> = samples_vec.iter().map(Vec::as_slice).collect();
208 let Some(dict) = train_zstd(0, &samples, 4096) else {
209 return; };
211 let plaintext = b"function test_case_99() { return 99; }\n";
212 let compressed = dict.compress(plaintext).expect("compress");
213 let recovered = dict
214 .decompress(&compressed, plaintext.len() as u32)
215 .expect("decompress");
216 assert_eq!(recovered.as_slice(), &plaintext[..]);
217 }
218
219 #[test]
220 fn allocate_ids_assigns_sequential_ids() {
221 let names = vec!["text", "binary", "source"];
222 let allocated = allocate_ids(&names).expect("allocate");
223 assert_eq!(allocated.len(), 3);
224 assert_eq!(allocated[0], ("text", 0));
225 assert_eq!(allocated[1], ("binary", 1));
226 assert_eq!(allocated[2], ("source", 2));
227 }
228
229 #[test]
230 fn allocate_ids_rejects_more_than_254_classes() {
231 let names: Vec<&str> = (0..255).map(|_| "x").collect();
232 assert!(allocate_ids(&names).is_err());
233 }
234}