1use alloc::vec::Vec;
20
21use crate::entropy::{
22 EncoderVariantForS, EntropyDecoder, EntropyEncoder, EntropyError, RawEncoder,
23};
24use crate::source::SliceSource;
25use crate::{Rans64Decoder, RansByteDecoder};
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq)]
29pub enum RansVariant {
30 RansByte,
32 Rans64,
34}
35
36impl RansVariant {
37 pub const fn as_int(&self) -> i32 {
39 match self {
40 RansVariant::RansByte => 1,
41 RansVariant::Rans64 => 0,
42 }
43 }
44
45 pub const fn from_int(v: i32) -> Option<Self> {
47 match v {
48 1 => Some(RansVariant::RansByte),
49 0 => Some(RansVariant::Rans64),
50 _ => None,
51 }
52 }
53}
54
55#[derive(Debug)]
61pub struct RansEncoderStream<S: EncoderVariantForS> {
62 encoder: Option<S::RawEnc>,
63 _s: core::marker::PhantomData<S>,
64}
65
66impl<S: EncoderVariantForS> Default for RansEncoderStream<S> {
67 fn default() -> Self {
68 Self::new()
69 }
70}
71
72impl<S: EncoderVariantForS> RansEncoderStream<S> {
73 pub fn new() -> Self {
75 Self {
76 encoder: None,
77 _s: core::marker::PhantomData,
78 }
79 }
80
81 pub fn is_initialized(&self) -> bool {
83 self.encoder.is_some()
84 }
85
86 pub fn push(
92 &mut self,
93 encoder: &EntropyEncoder<S>,
94 indices: &[i32],
95 values: &[i32],
96 ) -> Result<(), EntropyError> {
97 let raw = self.encoder_mut();
98 encoder.encode_batch(indices, values, raw)
99 }
100
101 pub fn flush(&mut self) -> Result<Vec<u8>, EntropyError> {
107 let mut raw = self.encoder.take().ok_or(EntropyError::InvalidState)?;
108 raw.flush();
109 let units = raw.into_units();
110 Ok(S::units_to_bytes(units))
111 }
112
113 pub fn reset(&mut self) {
115 self.encoder = None;
116 }
117
118 fn encoder_mut(&mut self) -> &mut S::RawEnc {
119 if self.encoder.is_none() {
120 self.encoder = Some(S::make_encoder());
121 }
122 self.encoder.as_mut().expect("just set")
123 }
124}
125
126#[derive(Debug)]
133pub struct RansDecoderStream<S: EncoderVariantForS> {
134 data: Vec<u8>,
135 cursor: Option<(usize, u64)>,
137 _s: core::marker::PhantomData<S>,
138}
139
140impl<S: EncoderVariantForS> Default for RansDecoderStream<S> {
141 fn default() -> Self {
142 Self::new()
143 }
144}
145
146impl<S: EncoderVariantForS> RansDecoderStream<S> {
147 pub fn new() -> Self {
149 Self {
150 data: Vec::new(),
151 cursor: None,
152 _s: core::marker::PhantomData,
153 }
154 }
155
156 pub fn open_on(data: &[u8]) -> Self {
158 Self {
159 data: data.to_vec(),
160 cursor: None,
161 _s: core::marker::PhantomData,
162 }
163 }
164
165 pub fn is_open(&self) -> bool {
167 !self.data.is_empty()
168 }
169
170 pub fn check_eof(&self) -> bool {
175 let Some((pos, state)) = self.cursor else {
176 return !self.is_open();
177 };
178 let unit_len = match S::NAME {
179 "RansByte" => self.data.len(),
180 "Rans64" => self.data.len() / 4,
181 _ => 0,
182 };
183 let lower = match S::NAME {
184 "RansByte" => 1u64 << 23,
185 "Rans64" => 1u64 << 31,
186 _ => 0,
187 };
188 pos == unit_len && state == lower
189 }
190
191 pub fn open(&mut self, data: &[u8]) {
193 self.data = data.to_vec();
194 self.cursor = None;
195 }
196
197 pub fn close(&mut self) {
199 self.data.clear();
200 self.cursor = None;
201 }
202
203 pub fn decode_eof(&mut self) -> Result<(), EntropyError> {
205 if !self.check_eof() {
206 return Err(EntropyError::InvalidStream);
207 }
208 self.close();
209 Ok(())
210 }
211
212 pub fn decode(
217 &mut self,
218 decoder: &EntropyDecoder<S>,
219 values: &mut [i32],
220 indices: &[i32],
221 ) -> Result<(), EntropyError> {
222 match S::NAME {
223 "RansByte" => {
224 let units = self.data.clone();
225 let mut source = SliceSource::new(&units);
226 let mut raw = match self.cursor {
227 Some((pos, state)) => {
228 source.seek(pos);
229 RansByteDecoder::from_state(source, state as u32)
230 }
231 None => {
232 let mut d = RansByteDecoder::new(source);
233 if !d.init() {
234 return Err(EntropyError::InvalidStream);
235 }
236 d
237 }
238 };
239 decoder.decode_byte_continue(&mut raw, values, indices)?;
240 self.cursor = Some((raw.source().position(), raw.state() as u64));
241 Ok(())
242 }
243 "Rans64" => {
244 if self.data.len() % 4 != 0 {
245 return Err(EntropyError::InvalidStream);
246 }
247 let units: Vec<u32> = self
248 .data
249 .chunks_exact(4)
250 .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
251 .collect();
252 let mut source = SliceSource::new(&units);
253 let mut raw = match self.cursor {
254 Some((pos, state)) => {
255 source.seek(pos);
256 Rans64Decoder::from_state(source, state)
257 }
258 None => {
259 let mut d = Rans64Decoder::new(source);
260 if !d.init() {
261 return Err(EntropyError::InvalidStream);
262 }
263 d
264 }
265 };
266 decoder.decode_64_continue(&mut raw, values, indices)?;
267 self.cursor = Some((raw.source().position(), raw.state()));
268 Ok(())
269 }
270 _ => Err(EntropyError::InvalidParams),
271 }
272 }
273
274 pub fn bytes_consumed(&self) -> usize {
276 match self.cursor {
277 Some((pos, _)) => match S::NAME {
278 "RansByte" => pos,
279 "Rans64" => pos * 4,
280 _ => 0,
281 },
282 None => 0,
283 }
284 }
285
286 pub fn data(&self) -> &[u8] {
288 &self.data
289 }
290}
291
292pub fn units_to_le_bytes(units: &[u32]) -> Vec<u8> {
294 let mut bytes = Vec::with_capacity(units.len() * 4);
295 for &u in units {
296 bytes.extend_from_slice(&u.to_le_bytes());
297 }
298 bytes
299}
300
301#[cfg(test)]
302mod tests {
303 use super::*;
304 use crate::entropy::EntropyDecoder;
305 use crate::variant::{Rans64, RansByte};
306
307 #[test]
308 fn test_variant_values() {
309 assert_eq!(RansVariant::RansByte.as_int(), 1);
310 assert_eq!(RansVariant::Rans64.as_int(), 0);
311 assert_eq!(RansVariant::from_int(1), Some(RansVariant::RansByte));
312 assert_eq!(RansVariant::from_int(0), Some(RansVariant::Rans64));
313 assert_eq!(RansVariant::from_int(2), None);
314 }
315
316 #[test]
317 fn test_encoder_stream_multipart() {
318 let pmf_lengths1 = vec![4, 6];
320 let pmf_offsets1 = vec![1, 2];
321 let pmf_table1 = vec![1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
322 let values1 = vec![-2, 1, 0, 1];
323 let indices1 = vec![0, 1, 0, 1];
324
325 let pmf_lengths2 = vec![5];
326 let pmf_offsets2 = vec![1];
327 let pmf_table2 = vec![1, 3, 3, 1, 1];
328 let values2 = vec![-2, 1, 2];
329 let indices2 = vec![0, 0, 0];
330
331 let mut encoder1 = EntropyEncoder::<RansByte>::new();
332 encoder1
333 .initialize(&pmf_lengths1, &pmf_offsets1, &pmf_table1, 16, 4)
334 .expect("init1");
335 let mut encoder2 = EntropyEncoder::<RansByte>::new();
336 encoder2
337 .initialize(&pmf_lengths2, &pmf_offsets2, &pmf_table2, 16, 4)
338 .expect("init2");
339
340 let mut stream = RansEncoderStream::<RansByte>::new();
341 stream.push(&encoder2, &indices2, &values2).expect("push2");
342 stream.push(&encoder1, &indices1, &values1).expect("push1");
343 let data = stream.flush().expect("flush");
344 assert!(!data.is_empty());
345
346 let mut decoder1 = EntropyDecoder::<RansByte>::new();
348 decoder1
349 .initialize(&pmf_lengths1, &pmf_offsets1, &pmf_table1, 16, 4)
350 .expect("dec1 init");
351 let mut decoder2 = EntropyDecoder::<RansByte>::new();
352 decoder2
353 .initialize(&pmf_lengths2, &pmf_offsets2, &pmf_table2, 16, 4)
354 .expect("dec2 init");
355
356 let mut dstream = RansDecoderStream::<RansByte>::open_on(&data);
357
358 let mut decoded1 = vec![0i32; values1.len()];
359 dstream
360 .decode(&decoder1, &mut decoded1, &indices1)
361 .expect("decode1");
362 assert_eq!(decoded1, values1);
363
364 let mut decoded2 = vec![0i32; values2.len()];
365 dstream
366 .decode(&decoder2, &mut decoded2, &indices2)
367 .expect("decode2");
368 assert_eq!(decoded2, values2);
369
370 dstream.decode_eof().expect("eof");
371 }
372
373 #[test]
374 fn test_encoder_stream_multipart_64() {
375 let pmf_lengths1 = vec![4, 6];
377 let pmf_offsets1 = vec![1, 2];
378 let pmf_table1 = vec![1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
379 let values1 = vec![-2, 1, 0, 1];
380 let indices1 = vec![0, 1, 0, 1];
381
382 let pmf_lengths2 = vec![5];
383 let pmf_offsets2 = vec![1];
384 let pmf_table2 = vec![1, 3, 3, 1, 1];
385 let values2 = vec![-2, 1, 2];
386 let indices2 = vec![0, 0, 0];
387
388 let mut encoder1 = EntropyEncoder::<Rans64>::new();
389 encoder1
390 .initialize(&pmf_lengths1, &pmf_offsets1, &pmf_table1, 16, 4)
391 .expect("init1");
392 let mut encoder2 = EntropyEncoder::<Rans64>::new();
393 encoder2
394 .initialize(&pmf_lengths2, &pmf_offsets2, &pmf_table2, 16, 4)
395 .expect("init2");
396
397 let mut stream = RansEncoderStream::<Rans64>::new();
398 stream.push(&encoder2, &indices2, &values2).expect("push2");
399 stream.push(&encoder1, &indices1, &values1).expect("push1");
400 let data = stream.flush().expect("flush");
401 assert!(!data.is_empty());
402 assert_eq!(data.len() % 4, 0, "Rans64 stream must be 4-byte aligned");
403
404 let mut decoder1 = EntropyDecoder::<Rans64>::new();
405 decoder1
406 .initialize(&pmf_lengths1, &pmf_offsets1, &pmf_table1, 16, 4)
407 .expect("dec1 init");
408 let mut decoder2 = EntropyDecoder::<Rans64>::new();
409 decoder2
410 .initialize(&pmf_lengths2, &pmf_offsets2, &pmf_table2, 16, 4)
411 .expect("dec2 init");
412
413 let mut dstream = RansDecoderStream::<Rans64>::open_on(&data);
414
415 let mut decoded1 = vec![0i32; values1.len()];
416 dstream
417 .decode(&decoder1, &mut decoded1, &indices1)
418 .expect("decode1");
419 assert_eq!(decoded1, values1);
420
421 let mut decoded2 = vec![0i32; values2.len()];
422 dstream
423 .decode(&decoder2, &mut decoded2, &indices2)
424 .expect("decode2");
425 assert_eq!(decoded2, values2);
426
427 dstream.decode_eof().expect("eof");
428 }
429
430 #[test]
431 fn test_encoder_stream_reuse_after_flush() {
432 let pmf_lengths = vec![4, 6];
433 let pmf_offsets = vec![1, 2];
434 let pmf_table = vec![1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
435 let values = vec![-2, 1, 0, 1];
436 let indices = vec![0, 1, 0, 1];
437
438 let mut encoder = EntropyEncoder::<RansByte>::new();
439 encoder
440 .initialize(&pmf_lengths, &pmf_offsets, &pmf_table, 16, 4)
441 .expect("init");
442
443 let mut stream = RansEncoderStream::<RansByte>::new();
444 stream.push(&encoder, &indices, &values).expect("push1");
445 let data1 = stream.flush().expect("flush1");
446
447 stream.push(&encoder, &indices, &values).expect("push2");
449 let data2 = stream.flush().expect("flush2");
450
451 let mut decoder = EntropyDecoder::<RansByte>::new();
452 decoder
453 .initialize(&pmf_lengths, &pmf_offsets, &pmf_table, 16, 4)
454 .expect("dec init");
455
456 for data in [&data1, &data2] {
457 let mut decoded = vec![0i32; values.len()];
458 decoder
459 .decode(&mut decoded, &indices, data)
460 .expect("decode");
461 assert_eq!(decoded, values);
462 }
463 }
464
465 #[test]
466 fn test_encoder_stream_reset_aborts() {
467 let pmf_lengths = vec![4, 6];
468 let pmf_offsets = vec![1, 2];
469 let pmf_table = vec![1, 3, 1, 1, 1, 3, 5, 3, 1, 1];
470 let values = vec![-2, 1, 0, 1];
471 let indices = vec![0, 1, 0, 1];
472
473 let mut encoder = EntropyEncoder::<RansByte>::new();
474 encoder
475 .initialize(&pmf_lengths, &pmf_offsets, &pmf_table, 16, 4)
476 .expect("init");
477
478 let mut stream = RansEncoderStream::<RansByte>::new();
479 stream.push(&encoder, &indices, &values).expect("push");
480 stream.reset();
481 assert!(!stream.is_initialized(), "reset must clear state");
482 assert!(stream.flush().is_err());
484 }
485
486 #[test]
487 fn test_decoder_stream_lifecycle() {
488 let mut stream = RansDecoderStream::<RansByte>::new();
489 assert!(!stream.is_open());
490 stream.open(&[1, 2, 3, 4, 5]);
491 assert!(stream.is_open());
492 assert!(!stream.check_eof());
493 stream.close();
494 assert!(!stream.is_open());
495 assert!(stream.check_eof());
496 }
497
498 #[test]
499 fn test_units_to_le_bytes() {
500 let units = [0x01020304u32, 0x05060708];
501 let bytes = <Rans64 as EncoderVariantForS>::units_to_bytes(units.to_vec());
502 assert_eq!(bytes, vec![4, 3, 2, 1, 8, 7, 6, 5]);
503 }
504}