Skip to main content

tuwunel_core/matrix/
event_id.rs

1//! Binary codec and fast hashing for event IDs.
2//!
3//! Since room v3 an event ID is the event's sha256 reference hash: `$` followed
4//! by 43 characters of unpadded base64 (standard alphabet in v3, URL-safe in
5//! v4+), so the 32 hash bytes and the ID convert losslessly in both directions.
6//!
7//! TODO: Move this into Ruma and generalize for other identifiers.
8
9use 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
19/// The sha256 reference hash encoded by a v3+ event ID.
20pub use crate::utils::hash::sha256::Digest as Sha256;
21use crate::utils::math::u64_from_usize_saturating;
22
23/// `BuildHasher` for event-ID-keyed maps and sets.
24///
25/// Uniformly distributed bits are recovered from the ID's base64 payload
26/// instead of running a byte hasher over the 44-byte string; other keys fall
27/// back to a folded-multiply hash. Seeded once per process; prefer `SipHash`
28/// where an attacker who can grind key bits must not learn bucket placement.
29#[derive(Clone, Copy, Debug, Default)]
30pub struct RandomState;
31
32/// Streaming state for [`RandomState`].
33#[derive(Clone, Copy, Debug, Default)]
34pub struct FoldHasher(u64);
35
36// The odd seed keeps the folded multiply bijective.
37static 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/// Decode an event ID into the sha256 it encodes. `None` unless the ID has the
54/// v3+ shape with canonical unpadded base64 in either alphabet.
55#[must_use]
56pub fn decode(event_id: &EventId) -> Option<Sha256> { decode_bytes(event_id.as_bytes()) }
57
58/// Encode a sha256 into the URL-safe (room v4+) form of the event ID. A v3
59/// (standard-alphabet) ID does not round-trip through this form.
60#[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	// The tail's three sextets carry the last two bytes; unpadded canonical
152	// base64 requires the two excess low bits to be zero.
153	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		// 'B' decodes to sextet 1, leaving a nonzero excess bit in the tail.
233		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}