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;
21
22/// `BuildHasher` for event-ID-keyed maps and sets.
23///
24/// Uniformly distributed bits are recovered from the ID's base64 payload
25/// instead of running a byte hasher over the 44-byte string; other keys fall
26/// back to a folded-multiply hash. Seeded once per process; prefer `SipHash`
27/// where an attacker who can grind key bits must not learn bucket placement.
28#[derive(Clone, Copy, Debug, Default)]
29pub struct RandomState;
30
31/// Streaming state for [`RandomState`].
32#[derive(Clone, Copy, Debug, Default)]
33pub struct FoldHasher(u64);
34
35// The odd seed keeps the folded multiply bijective.
36static 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/// Decode an event ID into the sha256 it encodes. `None` unless the ID has the
53/// v3+ shape with canonical unpadded base64 in either alphabet.
54#[must_use]
55pub fn decode(event_id: &EventId) -> Option<Sha256> { decode_bytes(event_id.as_bytes()) }
56
57/// Encode a sha256 into the URL-safe (room v4+) form of the event ID. A v3
58/// (standard-alphabet) ID does not round-trip through this form.
59#[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	// The tail's three sextets carry the last two bytes; unpadded canonical
151	// base64 requires the two excess low bits to be zero.
152	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		// 'B' decodes to sextet 1, leaving a nonzero excess bit in the tail.
232		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}