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