tuwunel_core/matrix/
event_id.rs1use std::{
10 hash::{BuildHasher, Hasher},
11 ops::BitOr,
12 sync::LazyLock,
13};
14
15use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
16use rand::random;
17use ruma::{EventId, OwnedEventId};
18
19pub use crate::utils::hash::sha256::Digest as Sha256;
21
22#[derive(Clone, Copy, Debug, Default)]
29pub struct RandomState;
30
31#[derive(Clone, Copy, Debug, Default)]
33pub struct FoldHasher(u64);
34
35static SEED: LazyLock<u64> = LazyLock::new(|| random::<u64>() | 1);
37
38static SEXTETS: [u8; SEXTETS_LEN] = sextets();
39
40const ALPHABET_STANDARD: &[u8; ALPHABET_LEN] =
41 b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
42
43const ALPHABET_URL_SAFE: &[u8; ALPHABET_LEN] =
44 b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_";
45
46const SEXTETS_LEN: usize = 256;
47const ALPHABET_LEN: usize = 64;
48const ENCODED_LEN: usize = 43;
49
50const INVALID: u8 = 0xFF;
51
52#[must_use]
55pub fn decode(event_id: &EventId) -> Option<Sha256> { decode_bytes(event_id.as_bytes()) }
56
57#[must_use]
60pub fn encode(sha256: &Sha256) -> OwnedEventId {
61 let mut encoded = [0_u8; ENCODED_LEN];
62 URL_SAFE_NO_PAD
63 .encode_slice(sha256, &mut encoded)
64 .expect("32 bytes always encode to 43 base64 characters");
65
66 let encoded: &str = str::from_utf8(&encoded).expect("base64 output is always ASCII");
67
68 OwnedEventId::from_parts('$', encoded, None).expect("valid event ID from base64 encoding")
69}
70
71impl BuildHasher for RandomState {
72 type Hasher = FoldHasher;
73
74 #[inline]
75 fn build_hasher(&self) -> FoldHasher { FoldHasher::default() }
76}
77
78impl Hasher for FoldHasher {
79 #[inline]
80 fn finish(&self) -> u64 { self.0 }
81
82 #[inline]
83 fn write(&mut self, bytes: &[u8]) { self.0 = fold(self.0 ^ hash_bytes(bytes)); }
84}
85
86fn hash_bytes(bytes: &[u8]) -> u64 {
87 decode_bytes(bytes).map_or_else(
88 || fold_bytes(bytes),
89 |sha256| {
90 sha256
91 .first_chunk()
92 .map(|first| u64::from_be_bytes(*first))
93 .expect("sha256 wider than u64")
94 },
95 )
96}
97
98fn fold_bytes(bytes: &[u8]) -> u64 {
99 let head = u64::try_from(bytes.len()).unwrap_or(u64::MAX);
100
101 bytes.chunks(8).fold(head, |acc, chunk| {
102 let mut word = [0_u8; 8];
103
104 word[..chunk.len()].copy_from_slice(chunk);
105 fold(acc ^ u64::from_be_bytes(word))
106 })
107}
108
109#[inline]
110fn fold(word: u64) -> u64 { folded_multiply(word, *SEED) }
111
112#[inline]
113#[expect(clippy::as_conversions, clippy::cast_possible_truncation)]
114fn folded_multiply(a: u64, b: u64) -> u64 {
115 let wide = u128::from(a).wrapping_mul(u128::from(b));
116
117 (wide as u64) ^ ((wide >> 64) as u64)
118}
119
120fn decode_bytes(bytes: &[u8]) -> Option<Sha256> {
121 bytes
122 .strip_prefix(b"$")
123 .and_then(|encoded| encoded.try_into().ok())
124 .and_then(decode_encoded)
125}
126
127fn decode_encoded(encoded: &[u8; ENCODED_LEN]) -> Option<Sha256> {
128 let sextet = |byte: &u8| u32::from(SEXTETS[usize::from(*byte)]);
129 let pack = |bytes: &[u8]| {
130 bytes
131 .iter()
132 .map(sextet)
133 .fold(0_u32, |group, sextet| (group << 6) | sextet)
134 };
135
136 let invalid = encoded
137 .iter()
138 .map(sextet)
139 .fold(0_u32, BitOr::bitor);
140
141 let mut sha256 = Sha256::default();
142 let (quads, tail) = encoded.as_chunks::<4>();
143 let (triples, _) = sha256.as_chunks_mut::<3>();
144
145 for (quad, bytes) in quads.iter().zip(triples) {
146 let [_, b0, b1, b2] = pack(quad).to_be_bytes();
147 *bytes = [b0, b1, b2];
148 }
149
150 let group = pack(tail);
153 let [_, _, b30, b31] = (group >> 2).to_be_bytes();
154
155 sha256[30] = b30;
156 sha256[31] = b31;
157 (invalid <= 0x3F && group.trailing_zeros() >= 2).then_some(sha256)
158}
159
160#[expect(clippy::as_conversions)]
161const fn sextets() -> [u8; SEXTETS_LEN] {
162 let mut table = [INVALID; SEXTETS_LEN];
163 let mut value = 0_u8;
164 while value < 64 {
165 table[ALPHABET_STANDARD[value as usize] as usize] = value;
166 table[ALPHABET_URL_SAFE[value as usize] as usize] = value;
167 value = value.wrapping_add(1);
168 }
169
170 table
171}
172
173#[cfg(test)]
174mod tests {
175 use std::{
176 collections::{HashMap, HashSet},
177 iter::repeat_with,
178 };
179
180 use base64::engine::general_purpose::STANDARD_NO_PAD;
181
182 use super::*;
183 use crate::utils::rand::event_id as random_event_id;
184
185 #[test]
186 fn roundtrip_random() {
187 for _ in 0..64 {
188 let event_id = random_event_id();
189 let sha256 = decode(&event_id).expect("random v4 event ID decodes");
190
191 assert_eq!(encode(&sha256), event_id);
192 }
193 }
194
195 #[test]
196 fn decode_standard_alphabet() {
197 for _ in 0..64 {
198 let sha256 = decode(&random_event_id()).unwrap();
199
200 let mut encoded = String::from("$");
201 STANDARD_NO_PAD.encode_string(sha256, &mut encoded);
202
203 let standard: OwnedEventId = encoded.try_into().unwrap();
204
205 assert_eq!(decode(&standard), Some(sha256));
206 }
207 }
208
209 #[test]
210 fn decode_rejects_non_v3_shapes() {
211 let legacy: OwnedEventId = "$legacy_event:server.example".try_into().unwrap();
212
213 assert_eq!(decode(&legacy), None);
214
215 let event_id = random_event_id();
216 let short: OwnedEventId = event_id
217 .as_str()
218 .get(..40)
219 .unwrap()
220 .try_into()
221 .unwrap();
222
223 assert_eq!(decode(&short), None);
224 }
225
226 #[test]
227 fn decode_rejects_non_canonical_tail() {
228 let event_id = random_event_id();
229 let mut noncanonical = event_id.as_str().to_owned();
230
231 noncanonical.replace_range(43.., "B");
233 let noncanonical: OwnedEventId = noncanonical.try_into().unwrap();
234
235 assert_eq!(decode(&noncanonical), None);
236 }
237
238 #[test]
239 fn map_operations() {
240 let mut map: HashMap<OwnedEventId, usize, RandomState> = HashMap::default();
241 let mut set: HashSet<OwnedEventId, RandomState> = HashSet::default();
242
243 let ids: Vec<OwnedEventId> = repeat_with(random_event_id)
244 .take(256)
245 .chain(["$legacy_event:server.example".try_into().unwrap()])
246 .collect();
247
248 for (i, event_id) in ids.iter().enumerate() {
249 map.insert(event_id.clone(), i);
250 set.insert(event_id.clone());
251 }
252
253 for (i, event_id) in ids.iter().enumerate() {
254 let event_id: &EventId = event_id;
255
256 assert_eq!(map.get(event_id), Some(&i));
257 assert!(set.contains(event_id));
258 }
259
260 assert_eq!(map.len(), ids.len());
261 }
262}