Skip to main content

safebrowsing_hash/
rice.rs

1//! Rice-Golomb encoding and decoding for Safe Browsing
2//!
3//! This module implements the Rice-Golomb encoding scheme used by the Safe Browsing API
4//! for compressing hash prefixes and indices.
5
6use crate::{HashError, Result};
7
8// Helper function for hex decoding in tests
9#[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// Wrapper for compatibility with test
18#[cfg(test)]
19mod hex {
20    pub fn decode(s: &str) -> Result<Vec<u8>, &'static str> {
21        Ok(super::hex_decode(s))
22    }
23}
24
25/// Bit reader for reading individual bits from a byte stream
26pub struct BitReader<'a> {
27    buf: &'a [u8],
28    mask: u8,
29}
30
31impl<'a> BitReader<'a> {
32    /// Create a new BitReader from byte data
33    pub fn new(buf: &'a [u8]) -> Self {
34        Self { buf, mask: 0x01 }
35    }
36
37    /// Read the specified number of bits and return as u32
38    /// Bits are read in little-endian order within each byte
39    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    /// Get the number of bits remaining to be read
73    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
84/// Rice-Golomb decoder for the Safe Browsing API
85///
86/// In Rice encoding, each number n is encoded as q and r where n = (q << k) + r.
87/// k is the Rice parameter (0..32).
88///
89/// The quotient q is encoded in unary: a sequence of q ones followed by a zero.
90/// The remainder r is encoded as a k-bit unsigned integer.
91pub struct RiceDecoder {
92    k: u32,
93}
94
95impl RiceDecoder {
96    /// Create a new Rice decoder with the given parameter k
97    pub fn new(k: u32) -> Self {
98        Self { k }
99    }
100
101    /// Read and decode a single value from the bit stream
102    pub fn read_value(&self, bit_reader: &mut BitReader) -> Result<u32> {
103        // Read quotient (unary encoding: count ones until we hit a zero)
104        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        // Read remainder (k bits)
114        let remainder = if self.k > 0 {
115            bit_reader.read_bits(self.k)?
116        } else {
117            0
118        };
119
120        // Combine quotient and remainder: n = (q << k) + r
121        Ok((quotient << self.k) + remainder)
122    }
123}
124
125/// Decode Rice-Golomb encoded integers from Safe Browsing format
126pub 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    // Start with the first value
145    let mut values = vec![first_value as u32];
146
147    // If no additional entries, just return the first value
148    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    // Decode each delta and add to the running sum
156    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    // Check that we consumed most of the input (allow up to 7 unused bits)
163    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        // Test data: 0b10110100, 0b11010011
179        let data = vec![0b10110100, 0b11010011];
180        let mut reader = BitReader::new(&data);
181
182        // Read individual bits (LSB first within each byte)
183        assert_eq!(reader.read_bits(1).unwrap(), 0); // bit 0 of first byte
184        assert_eq!(reader.read_bits(1).unwrap(), 0); // bit 1 of first byte
185        assert_eq!(reader.read_bits(1).unwrap(), 1); // bit 2 of first byte
186        assert_eq!(reader.read_bits(1).unwrap(), 0); // bit 3 of first byte
187
188        // Read multiple bits at once
189        assert_eq!(reader.read_bits(4).unwrap(), 0b1011); // bits 4-7 of first byte
190
191        // Cross byte boundary
192        assert_eq!(reader.read_bits(4).unwrap(), 0b0011); // bits 0-3 of second byte
193        assert_eq!(reader.read_bits(4).unwrap(), 0b1101); // bits 4-7 of second byte
194    }
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        // Test vectors from Go implementation
222        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        // Test complete rice integer decoding with Go test vectors
265        let test_cases = vec![
266            (2, 0, 2, "f702", vec![0, 15, 24]), // first_value=0, num_entries=2
267            (5, 42, 1, "00", vec![42, 42]),     // first_value=42, num_entries=1
268        ];
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}