1use crate::{HashError, Result};
7
8#[cfg(test)]
10fn hex_decode(s: &str) -> Vec<u8> {
11 (0..s.len())
12 .step_by(2)
13 .map(|i| u8::from_str_radix(&s[i..i + 2], 16).unwrap())
14 .collect()
15}
16
17#[cfg(test)]
19mod hex {
20 pub fn decode(s: &str) -> Result<Vec<u8>, &'static str> {
21 Ok(super::hex_decode(s))
22 }
23}
24
25pub struct BitReader<'a> {
27 buf: &'a [u8],
28 mask: u8,
29}
30
31impl<'a> BitReader<'a> {
32 pub fn new(buf: &'a [u8]) -> Self {
34 Self { buf, mask: 0x01 }
35 }
36
37 pub fn read_bits(&mut self, num_bits: u32) -> Result<u32> {
40 if num_bits == 0 {
41 return Ok(0);
42 }
43 if num_bits > 32 {
44 return Err(HashError::InvalidFormat(
45 "Cannot read more than 32 bits".to_string(),
46 ));
47 }
48
49 let mut result = 0u32;
50
51 for i in 0..num_bits {
52 if self.buf.is_empty() {
53 return Err(HashError::InvalidFormat(
54 "Unexpected end of data".to_string(),
55 ));
56 }
57
58 if (self.buf[0] & self.mask) > 0 {
59 result |= 1u32 << i;
60 }
61
62 self.mask <<= 1;
63 if self.mask == 0 {
64 self.buf = &self.buf[1..];
65 self.mask = 0x01;
66 }
67 }
68
69 Ok(result)
70 }
71
72 pub fn bits_remaining(&self) -> usize {
74 let mut n = 8 * self.buf.len();
75 let mut m = self.mask | 1;
76 while m != 1 {
77 n -= 1;
78 m >>= 1;
79 }
80 n
81 }
82}
83
84pub struct RiceDecoder {
92 k: u32,
93}
94
95impl RiceDecoder {
96 pub fn new(k: u32) -> Self {
98 Self { k }
99 }
100
101 pub fn read_value(&self, bit_reader: &mut BitReader) -> Result<u32> {
103 let mut quotient = 0u32;
105 loop {
106 let bit = bit_reader.read_bits(1)?;
107 if bit == 0 {
108 break;
109 }
110 quotient += 1;
111 }
112
113 let remainder = if self.k > 0 {
115 bit_reader.read_bits(self.k)?
116 } else {
117 0
118 };
119
120 Ok((quotient << self.k) + remainder)
122 }
123}
124
125pub fn decode_rice_integers(
127 rice_parameter: i32,
128 first_value: i64,
129 num_entries: i32,
130 encoded_data: &[u8],
131) -> Result<Vec<u32>> {
132 if !(0..=32).contains(&rice_parameter) {
133 return Err(HashError::InvalidFormat(format!(
134 "Invalid rice parameter: {rice_parameter}"
135 )));
136 }
137
138 if num_entries < 0 {
139 return Err(HashError::InvalidFormat(format!(
140 "Invalid num_entries: {num_entries}"
141 )));
142 }
143
144 let mut values = vec![first_value as u32];
146
147 if num_entries == 0 {
149 return Ok(values);
150 }
151
152 let mut bit_reader = BitReader::new(encoded_data);
153 let rice_decoder = RiceDecoder::new(rice_parameter as u32);
154
155 for i in 0..num_entries {
157 let delta = rice_decoder.read_value(&mut bit_reader)?;
158 let next_value = values[i as usize] + delta;
159 values.push(next_value);
160 }
161
162 if bit_reader.bits_remaining() >= 8 {
164 return Err(HashError::InvalidFormat(
165 "Unconsumed rice encoded data".to_string(),
166 ));
167 }
168
169 Ok(values)
170}
171
172#[cfg(test)]
173mod tests {
174 use super::*;
175
176 #[test]
177 fn test_bit_reader_basic() {
178 let data = vec![0b10110100, 0b11010011];
180 let mut reader = BitReader::new(&data);
181
182 assert_eq!(reader.read_bits(1).unwrap(), 0); assert_eq!(reader.read_bits(1).unwrap(), 0); assert_eq!(reader.read_bits(1).unwrap(), 1); assert_eq!(reader.read_bits(1).unwrap(), 0); assert_eq!(reader.read_bits(4).unwrap(), 0b1011); assert_eq!(reader.read_bits(4).unwrap(), 0b0011); assert_eq!(reader.read_bits(4).unwrap(), 0b1101); }
195
196 #[test]
197 fn test_bit_reader_empty() {
198 let data = vec![];
199 let mut reader = BitReader::new(&data);
200
201 assert!(reader.read_bits(1).is_err());
202 assert_eq!(reader.bits_remaining(), 0);
203 }
204
205 #[test]
206 fn test_bit_reader_bits_remaining() {
207 let data = vec![0xFF, 0xFF];
208 let mut reader = BitReader::new(&data);
209
210 assert_eq!(reader.bits_remaining(), 16);
211 reader.read_bits(3).unwrap();
212 assert_eq!(reader.bits_remaining(), 13);
213 reader.read_bits(8).unwrap();
214 assert_eq!(reader.bits_remaining(), 5);
215 reader.read_bits(5).unwrap();
216 assert_eq!(reader.bits_remaining(), 0);
217 }
218
219 #[test]
220 fn test_rice_decoder_with_go_vectors() {
221 let test_cases = vec![
223 (2, "f702", vec![15, 9]),
224 (5, "00", vec![0]),
225 (
226 28,
227 "54607be70a5fc1dcee69defe583ca3d6a5f2108c4a595600",
228 vec![
229 62763050, 1046523781, 192522171, 1800511020, 4442775, 582142548,
230 ],
231 ),
232 ];
233
234 for (k, hex_input, expected) in test_cases {
235 let data = hex::decode(hex_input).unwrap();
236 let mut bit_reader = BitReader::new(&data);
237 let decoder = RiceDecoder::new(k);
238
239 let mut results = Vec::new();
240 for _ in 0..expected.len() {
241 let value = decoder.read_value(&mut bit_reader).unwrap();
242 results.push(value);
243 }
244
245 assert_eq!(results, expected, "Failed for k={k}, input={hex_input}");
246 }
247 }
248
249 #[test]
250 fn test_decode_rice_integers_empty() {
251 let result = decode_rice_integers(2, 42, 0, &[]).unwrap();
252 assert_eq!(result, vec![42]);
253 }
254
255 #[test]
256 fn test_decode_rice_integers_invalid_params() {
257 assert!(decode_rice_integers(-1, 0, 1, &[0]).is_err());
258 assert!(decode_rice_integers(33, 0, 1, &[0]).is_err());
259 assert!(decode_rice_integers(2, 0, -1, &[0]).is_err());
260 }
261
262 #[test]
263 fn test_decode_rice_integers_with_go_vectors() {
264 let test_cases = vec![
266 (2, 0, 2, "f702", vec![0, 15, 24]), (5, 42, 1, "00", vec![42, 42]), ];
269
270 for (k, first_value, num_entries, hex_input, expected) in test_cases {
271 let data = hex::decode(hex_input).unwrap();
272 let result = decode_rice_integers(k, first_value, num_entries, &data).unwrap();
273 assert_eq!(
274 result, expected,
275 "Failed for k={k}, first={first_value}, entries={num_entries}, input={hex_input}"
276 );
277 }
278 }
279}