Skip to main content

ahu/
index.rs

1use byteorder::{BigEndian, ReadBytesExt, WriteBytesExt};
2use sha2::{Digest, Sha256};
3use std::io::{Read, Write};
4
5use crate::error::{AhuError, Result};
6
7pub const INDEX_RECORD_SIZE: usize = 48;
8
9bitflags::bitflags! {
10    #[derive(Debug, Clone, Copy, PartialEq, Eq)]
11    pub struct IndexFlags: u16 {
12        const MULTI     = 0b0000_0000_0000_0001;
13        const ALIAS     = 0b0000_0000_0000_0010;
14        const TOMBSTONE = 0b0000_0000_0000_0100;
15    }
16}
17
18/// Algorithm discriminator for dual-algorithm bundles.
19///
20/// In single-algorithm bundles, all records have discriminator 0 (default).
21/// In dual-algorithm bundles, the classical response uses 0 and the
22/// post-quantum variant uses the algorithm-specific value.
23pub const ALG_DISC_DEFAULT: u16 = 0;
24pub const ALG_DISC_ML_DSA_44: u16 = 2;
25pub const ALG_DISC_ML_DSA_65: u16 = 3;
26pub const ALG_DISC_ML_DSA_87: u16 = 4;
27
28#[derive(Debug, Clone, PartialEq, Eq)]
29pub struct IndexRecord {
30    pub entry_key: [u8; 32],
31    pub data_offset: u64,
32    pub data_length: u32,
33    pub flags: IndexFlags,
34    pub discriminator: u16,
35}
36
37impl IndexRecord {
38    pub fn read_from<R: Read>(reader: &mut R) -> Result<Self> {
39        let mut entry_key = [0u8; 32];
40        reader.read_exact(&mut entry_key)?;
41        let data_offset = reader.read_u64::<BigEndian>()?;
42        let data_length = reader.read_u32::<BigEndian>()?;
43        let flags_raw = reader.read_u16::<BigEndian>()?;
44        let discriminator = reader.read_u16::<BigEndian>()?;
45        let flags = IndexFlags::from_bits_truncate(flags_raw);
46
47        Ok(IndexRecord {
48            entry_key,
49            data_offset,
50            data_length,
51            flags,
52            discriminator,
53        })
54    }
55
56    pub fn write_to<W: Write>(&self, writer: &mut W) -> Result<()> {
57        writer.write_all(&self.entry_key)?;
58        writer.write_u64::<BigEndian>(self.data_offset)?;
59        writer.write_u32::<BigEndian>(self.data_length)?;
60        writer.write_u16::<BigEndian>(self.flags.bits())?;
61        writer.write_u16::<BigEndian>(self.discriminator)?;
62        Ok(())
63    }
64
65    pub fn is_tombstone(&self) -> bool {
66        self.flags.contains(IndexFlags::TOMBSTONE)
67    }
68
69    pub fn is_alias(&self) -> bool {
70        self.flags.contains(IndexFlags::ALIAS)
71    }
72
73    pub fn is_multi(&self) -> bool {
74        self.flags.contains(IndexFlags::MULTI)
75    }
76}
77
78pub fn compute_entry_key(certid_der: &[u8]) -> [u8; 32] {
79    let mut hasher = Sha256::new();
80    hasher.update(certid_der);
81    hasher.finalize().into()
82}
83
84pub fn validate_sort_order(records: &[IndexRecord]) -> Result<()> {
85    for i in 0..records.len().saturating_sub(1) {
86        let key_ord = records[i].entry_key.cmp(&records[i + 1].entry_key);
87        match key_ord {
88            std::cmp::Ordering::Greater => {
89                return Err(AhuError::IndexNotSorted {
90                    index: i,
91                    key: hex::encode(records[i].entry_key),
92                    next_key: hex::encode(records[i + 1].entry_key),
93                });
94            }
95            std::cmp::Ordering::Equal => {
96                match records[i].discriminator.cmp(&records[i + 1].discriminator) {
97                    std::cmp::Ordering::Equal => {
98                        return Err(AhuError::DuplicateEntryKey {
99                            index: i,
100                            key: hex::encode(records[i].entry_key),
101                        });
102                    }
103                    std::cmp::Ordering::Greater => {
104                        return Err(AhuError::IndexNotSorted {
105                            index: i,
106                            key: hex::encode(records[i].entry_key),
107                            next_key: hex::encode(records[i + 1].entry_key),
108                        });
109                    }
110                    std::cmp::Ordering::Less => {}
111                }
112            }
113            std::cmp::Ordering::Less => {}
114        }
115    }
116    Ok(())
117}
118
119/// Search for a record with the given key and discriminator 0 (default).
120pub fn binary_search(records: &[IndexRecord], entry_key: &[u8; 32]) -> Option<usize> {
121    binary_search_with_discriminator(records, entry_key, ALG_DISC_DEFAULT)
122}
123
124/// Search for a record with the given key and specific discriminator.
125pub fn binary_search_with_discriminator(
126    records: &[IndexRecord],
127    entry_key: &[u8; 32],
128    discriminator: u16,
129) -> Option<usize> {
130    records
131        .binary_search_by(|r| {
132            r.entry_key
133                .cmp(entry_key)
134                .then(r.discriminator.cmp(&discriminator))
135        })
136        .ok()
137}
138
139/// Search for the best-matching record given a list of preferred discriminators.
140/// Tries each discriminator in order; returns the first hit. Falls back to
141/// discriminator 0 if none of the preferences match.
142pub fn binary_search_preferred(
143    records: &[IndexRecord],
144    entry_key: &[u8; 32],
145    preferences: &[u16],
146) -> Option<usize> {
147    for &disc in preferences {
148        if let Some(idx) = binary_search_with_discriminator(records, entry_key, disc) {
149            return Some(idx);
150        }
151    }
152    if !preferences.contains(&ALG_DISC_DEFAULT) {
153        binary_search_with_discriminator(records, entry_key, ALG_DISC_DEFAULT)
154    } else {
155        None
156    }
157}
158
159#[cfg(test)]
160mod tests {
161    use super::*;
162    use std::io::Cursor;
163
164    fn rec(key_byte: u8, disc: u16) -> IndexRecord {
165        IndexRecord {
166            entry_key: {
167                let mut k = [0u8; 32];
168                k[0] = key_byte;
169                k
170            },
171            data_offset: 0,
172            data_length: 100,
173            flags: IndexFlags::empty(),
174            discriminator: disc,
175        }
176    }
177
178    #[test]
179    fn record_round_trip() {
180        let r = IndexRecord {
181            entry_key: [0xAB; 32],
182            data_offset: 1024,
183            data_length: 512,
184            flags: IndexFlags::MULTI | IndexFlags::ALIAS,
185            discriminator: 0,
186        };
187
188        let mut buf = Vec::new();
189        r.write_to(&mut buf).unwrap();
190        assert_eq!(buf.len(), INDEX_RECORD_SIZE);
191
192        let mut cursor = Cursor::new(&buf);
193        let read_back = IndexRecord::read_from(&mut cursor).unwrap();
194        assert_eq!(r, read_back);
195    }
196
197    #[test]
198    fn record_round_trip_with_discriminator() {
199        let r = IndexRecord {
200            entry_key: [0xCD; 32],
201            data_offset: 2048,
202            data_length: 256,
203            flags: IndexFlags::empty(),
204            discriminator: ALG_DISC_ML_DSA_87,
205        };
206
207        let mut buf = Vec::new();
208        r.write_to(&mut buf).unwrap();
209
210        let mut cursor = Cursor::new(&buf);
211        let read_back = IndexRecord::read_from(&mut cursor).unwrap();
212        assert_eq!(r, read_back);
213        assert_eq!(read_back.discriminator, ALG_DISC_ML_DSA_87);
214    }
215
216    #[test]
217    fn sort_order_validation() {
218        let r1 = rec(0x01, 0);
219        let r2 = rec(0x02, 0);
220
221        validate_sort_order(&[r1.clone(), r2.clone()]).unwrap();
222
223        let err = validate_sort_order(&[r2, r1]).unwrap_err();
224        assert!(matches!(err, AhuError::IndexNotSorted { .. }));
225    }
226
227    #[test]
228    fn sort_order_same_key_different_discriminators() {
229        let r1 = rec(0x01, 0);
230        let r2 = rec(0x01, ALG_DISC_ML_DSA_87);
231        validate_sort_order(&[r1, r2]).unwrap();
232    }
233
234    #[test]
235    fn sort_order_same_key_same_discriminator_rejected() {
236        let r1 = rec(0x01, ALG_DISC_ML_DSA_87);
237        let r2 = rec(0x01, ALG_DISC_ML_DSA_87);
238        let err = validate_sort_order(&[r1, r2]).unwrap_err();
239        assert!(matches!(err, AhuError::DuplicateEntryKey { .. }));
240    }
241
242    #[test]
243    fn sort_order_same_key_wrong_discriminator_order() {
244        let r1 = rec(0x01, ALG_DISC_ML_DSA_87);
245        let r2 = rec(0x01, 0);
246        let err = validate_sort_order(&[r1, r2]).unwrap_err();
247        assert!(matches!(err, AhuError::IndexNotSorted { .. }));
248    }
249
250    #[test]
251    fn binary_search_finds_key() {
252        let records: Vec<IndexRecord> = (0..10u8).map(|i| rec(i * 10, 0)).collect();
253
254        let mut target = [0u8; 32];
255        target[0] = 50;
256        assert_eq!(binary_search(&records, &target), Some(5));
257
258        target[0] = 55;
259        assert_eq!(binary_search(&records, &target), None);
260    }
261
262    #[test]
263    fn binary_search_with_discriminator_finds_variant() {
264        let records = vec![rec(0x01, 0), rec(0x01, ALG_DISC_ML_DSA_87), rec(0x02, 0)];
265
266        let mut key = [0u8; 32];
267        key[0] = 0x01;
268
269        assert_eq!(binary_search(&records, &key), Some(0));
270        assert_eq!(
271            binary_search_with_discriminator(&records, &key, ALG_DISC_ML_DSA_87),
272            Some(1)
273        );
274        assert_eq!(
275            binary_search_with_discriminator(&records, &key, ALG_DISC_ML_DSA_44),
276            None
277        );
278    }
279
280    #[test]
281    fn binary_search_preferred_tries_in_order() {
282        let records = vec![
283            rec(0x01, 0),
284            rec(0x01, ALG_DISC_ML_DSA_87),
285            rec(0x02, 0),
286            rec(0x02, ALG_DISC_ML_DSA_87),
287        ];
288
289        let mut key = [0u8; 32];
290        key[0] = 0x01;
291
292        assert_eq!(
293            binary_search_preferred(&records, &key, &[ALG_DISC_ML_DSA_87]),
294            Some(1)
295        );
296        assert_eq!(
297            binary_search_preferred(&records, &key, &[ALG_DISC_ML_DSA_44, ALG_DISC_ML_DSA_87]),
298            Some(1)
299        );
300        // Falls back to default when preference not found
301        assert_eq!(
302            binary_search_preferred(&records, &key, &[ALG_DISC_ML_DSA_44]),
303            Some(0)
304        );
305        // Empty preferences → default
306        assert_eq!(binary_search_preferred(&records, &key, &[]), Some(0));
307    }
308}