Skip to main content

tuwunel_service/rendezvous/
mod.rs

1// RAM fits 4 KiB, minute-lived sessions; restarts end them like OAuth state.
2
3use std::{
4	cmp::max,
5	collections::{BTreeMap, HashMap},
6	net::IpAddr,
7	str,
8	sync::{Arc, Mutex, RwLock},
9	time::{Duration, Instant, SystemTime, UNIX_EPOCH},
10};
11
12use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD as b64};
13use bytes::Bytes;
14use http::StatusCode;
15use ruma::api::error::{ErrorKind, LimitExceededErrorData};
16use tuwunel_core::{
17	Error, Result,
18	arrayvec::ArrayString,
19	implement,
20	utils::{hash::sha256::concat, rand::string_array, time::duration_since_epoch},
21};
22
23pub type SessionId = ArrayString<SESSION_ID_LENGTH>;
24pub type Etag = ArrayString<ETAG_LENGTH>;
25type Sessions = BTreeMap<SessionId, Session>;
26type Ratelimiter = Mutex<HashMap<IpAddr, (Instant, f64)>>;
27
28pub struct Service {
29	sessions: RwLock<Sessions>,
30	// At most 4096 short-lived per-IP buckets.
31	ratelimiter: Ratelimiter,
32	services: Arc<crate::services::OnceServices>,
33}
34
35struct Session {
36	data: Bytes,
37	etag: Etag,
38	created: SystemTime,
39	last_modified: SystemTime,
40	expires_at: SystemTime,
41}
42
43#[derive(Clone, Copy, Debug, Eq, PartialEq)]
44pub struct Meta {
45	pub etag: Etag,
46	pub expires_at: SystemTime,
47	pub last_modified: SystemTime,
48}
49
50#[derive(Clone, Debug, Eq, PartialEq)]
51pub enum Get {
52	Data {
53		data: Bytes,
54		meta: Meta,
55	},
56	NotModified(Meta),
57	NotFound,
58}
59
60#[derive(Clone, Copy, Debug, Eq, PartialEq)]
61pub enum Put {
62	Accepted(Meta),
63	PreconditionFailed(Meta),
64	NotFound,
65}
66
67#[derive(Clone, Copy)]
68enum Validator<'a> {
69	Etag(&'a str),
70	SequenceToken(&'a str),
71}
72
73const SESSION_ID_LENGTH: usize = 32;
74const ETAG_VALUE_LENGTH: usize = 43;
75const ETAG_LENGTH: usize = ETAG_VALUE_LENGTH + 2;
76const MILLIS_PER_SECOND: u64 = 1000;
77const MAX_HTTP_DATE_SECONDS: u64 = 253_402_300_799;
78const MONOTONIC_STEP: Duration = Duration::from_millis(1);
79const RATELIMIT_MAP_CAP: usize = 4096;
80
81#[cfg(test)]
82mod tests;
83
84impl crate::Service for Service {
85	fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
86		Ok(Arc::new(Self {
87			sessions: RwLock::new(Sessions::new()),
88			ratelimiter: Mutex::new(HashMap::new()),
89			services: args.services.clone(),
90		}))
91	}
92
93	fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
94}
95
96#[implement(Service)]
97pub fn check_rate_limit(&self, client: IpAddr) -> Result {
98	let config = &self.services.server.config;
99	let rate = f64::from(config.rendezvous_rc_per_second.max(1));
100	let burst = f64::from(config.rendezvous_rc_burst_count.max(1));
101
102	check_bucket_at(&self.ratelimiter, client, rate, burst, Instant::now())
103}
104
105fn check_bucket_at(
106	table: &Ratelimiter,
107	client: IpAddr,
108	rate: f64,
109	burst: f64,
110	now: Instant,
111) -> Result {
112	let mut buckets = table.lock()?;
113
114	if buckets.len() >= RATELIMIT_MAP_CAP && !buckets.contains_key(&client) {
115		let mut oldest = None;
116
117		buckets.retain(|client, (last, tokens)| {
118			let refilled = now
119				.duration_since(*last)
120				.as_secs_f64()
121				.mul_add(rate, *tokens);
122			let retain = refilled < burst;
123
124			if retain && oldest.is_none_or(|(_, oldest_at)| *last < oldest_at) {
125				oldest = Some((*client, *last));
126			}
127
128			retain
129		});
130
131		if buckets.len() >= RATELIMIT_MAP_CAP
132			&& let Some((oldest, _)) = oldest
133		{
134			buckets.remove(&oldest);
135		}
136	}
137
138	let (last_time, tokens) = buckets.entry(client).or_insert((now, burst));
139	let new_tokens = now
140		.duration_since(*last_time)
141		.as_secs_f64()
142		.mul_add(rate, *tokens)
143		.min(burst);
144
145	if new_tokens < 1.0 {
146		return Err(Error::Request(
147			ErrorKind::LimitExceeded(LimitExceededErrorData { retry_after: None }),
148			"Too many rendezvous requests.".into(),
149			StatusCode::TOO_MANY_REQUESTS,
150		));
151	}
152
153	*last_time = now;
154	*tokens = new_tokens - 1.0;
155
156	Ok(())
157}
158
159#[implement(Service)]
160pub fn create(&self, data: Bytes) -> (SessionId, Meta) {
161	let config = &self.services.server.config;
162	let ttl = Duration::from_secs(config.rendezvous_session_ttl);
163
164	self.create_at(data, SystemTime::now(), ttl, config.rendezvous_max_sessions)
165}
166
167#[implement(Service)]
168fn create_at(
169	&self,
170	data: Bytes,
171	now: SystemTime,
172	ttl: Duration,
173	max_sessions: usize,
174) -> (SessionId, Meta) {
175	let capacity = max_sessions.max(1);
176	let mut id = string_array::<SESSION_ID_LENGTH>();
177	let expires_at = expires_at(now, ttl);
178
179	let mut sessions = self.sessions.write().expect("locked for writing");
180
181	sessions.retain(|_, session| session.expires_at > now);
182	while sessions.contains_key(id.as_str()) {
183		id = string_array::<SESSION_ID_LENGTH>();
184	}
185
186	while sessions.len() >= capacity {
187		let Some(id) = sessions
188			.iter()
189			.min_by_key(|(_, session)| session.created)
190			.map(|(id, _)| *id)
191		else {
192			break;
193		};
194
195		sessions.remove(&id);
196	}
197
198	let session = Session {
199		etag: etag(id.as_str(), &data, now),
200		data,
201		created: now,
202		last_modified: now,
203		expires_at,
204	};
205
206	let meta = session.meta();
207
208	sessions.insert(id, session);
209
210	(id, meta)
211}
212
213#[implement(Service)]
214pub fn get(&self, id: &str, if_none_match: Option<&str>) -> Get {
215	self.get_at(id, if_none_match, SystemTime::now())
216}
217
218#[implement(Service)]
219fn get_at(&self, id: &str, if_none_match: Option<&str>, now: SystemTime) -> Get {
220	{
221		let sessions = self.sessions.read().expect("locked for reading");
222
223		match sessions.get(id) {
224			| Some(session) if session.expires_at > now => {
225				return get_outcome(session, if_none_match);
226			},
227			| Some(_) => {},
228			| None => return Get::NotFound,
229		}
230	}
231
232	let mut sessions = self.sessions.write().expect("locked for writing");
233	if sessions
234		.get(id)
235		.is_some_and(|session| session.expires_at <= now)
236	{
237		sessions.remove(id);
238
239		return Get::NotFound;
240	}
241
242	sessions
243		.get(id)
244		.map_or(Get::NotFound, |session| get_outcome(session, if_none_match))
245}
246
247#[implement(Service)]
248pub fn put(&self, id: &str, if_match: &str, data: Bytes) -> Put {
249	let ttl = Duration::from_secs(self.services.server.config.rendezvous_session_ttl);
250
251	self.put_at(id, if_match, data, SystemTime::now(), ttl)
252}
253
254#[implement(Service)]
255fn put_at(&self, id: &str, if_match: &str, data: Bytes, now: SystemTime, ttl: Duration) -> Put {
256	self.put_with_at(id, Validator::Etag(if_match), data, now, ttl)
257}
258
259#[implement(Service)]
260pub fn put_token(&self, id: &str, sequence_token: &str, data: Bytes) -> Put {
261	let ttl = Duration::from_secs(self.services.server.config.rendezvous_session_ttl);
262
263	self.put_token_at(id, sequence_token, data, SystemTime::now(), ttl)
264}
265
266#[implement(Service)]
267fn put_token_at(
268	&self,
269	id: &str,
270	sequence_token: &str,
271	data: Bytes,
272	now: SystemTime,
273	ttl: Duration,
274) -> Put {
275	self.put_with_at(id, Validator::SequenceToken(sequence_token), data, now, ttl)
276}
277
278#[implement(Service)]
279fn put_with_at(
280	&self,
281	id: &str,
282	validator: Validator<'_>,
283	data: Bytes,
284	now: SystemTime,
285	ttl: Duration,
286) -> Put {
287	let mut sessions = self.sessions.write().expect("locked for writing");
288
289	if sessions
290		.get(id)
291		.is_some_and(|session| session.expires_at <= now)
292	{
293		sessions.remove(id);
294
295		return Put::NotFound;
296	}
297
298	let Some(session) = sessions.get_mut(id) else {
299		return Put::NotFound;
300	};
301
302	if !validator.matches(&session.etag) {
303		if data != session.data {
304			return Put::PreconditionFailed(session.meta());
305		}
306
307		session.expires_at = expires_at(now, ttl);
308
309		return Put::Accepted(session.meta());
310	}
311
312	let created = session.created;
313	let last_modified = next_last_modified(now, session.last_modified);
314	let expires_at = expires_at(now, ttl);
315
316	*session = Session {
317		etag: etag(id, &data, last_modified),
318		data,
319		created,
320		last_modified,
321		expires_at,
322	};
323
324	Put::Accepted(session.meta())
325}
326
327#[implement(Service)]
328pub fn delete(&self, id: &str) -> bool {
329	self.sessions
330		.write()
331		.expect("locked for writing")
332		.remove(id)
333		.is_some()
334}
335
336#[implement(Service)]
337pub fn delete_if_active(&self, id: &str) -> bool {
338	self.delete_if_active_at(id, SystemTime::now())
339}
340
341#[implement(Service)]
342fn delete_if_active_at(&self, id: &str, now: SystemTime) -> bool {
343	self.sessions
344		.write()
345		.expect("locked for writing")
346		.remove(id)
347		.is_some_and(|session| session.expires_at > now)
348}
349
350#[implement(Meta)]
351#[must_use]
352#[inline]
353pub fn sequence_token(&self) -> &str { etag_value(&self.etag) }
354
355fn etag_value(etag: &Etag) -> &str {
356	etag.as_str()
357		.strip_prefix('"')
358		.and_then(|value| value.strip_suffix('"'))
359		.expect("ETag is quoted")
360}
361
362#[implement(Meta)]
363#[must_use]
364#[inline]
365pub fn expires_in(&self) -> Duration { self.expires_in_at(SystemTime::now()) }
366
367#[implement(Meta)]
368fn expires_in_at(&self, now: SystemTime) -> Duration {
369	self.expires_at
370		.duration_since(now)
371		.unwrap_or_default()
372}
373
374impl Session {
375	fn meta(&self) -> Meta {
376		Meta {
377			etag: self.etag,
378			expires_at: self.expires_at,
379			last_modified: self.last_modified,
380		}
381	}
382}
383
384impl Validator<'_> {
385	fn matches(self, stored: &Etag) -> bool {
386		match self {
387			| Self::Etag(candidate) => etag_matches(candidate, stored),
388			| Self::SequenceToken(candidate) => candidate == etag_value(stored),
389		}
390	}
391}
392
393fn expires_at(now: SystemTime, ttl: Duration) -> SystemTime {
394	let latest = UNIX_EPOCH
395		.checked_add(Duration::from_secs(MAX_HTTP_DATE_SECONDS))
396		.expect("latest HTTP date should fit in SystemTime");
397
398	now.checked_add(ttl)
399		.map_or(latest, |expires_at| expires_at.min(latest))
400}
401
402fn next_last_modified(now: SystemTime, previous: SystemTime) -> SystemTime {
403	previous
404		.checked_add(MONOTONIC_STEP)
405		.map_or(now, |next| max(now, next))
406}
407
408fn etag(id: &str, data: &Bytes, last_modified: SystemTime) -> Etag {
409	let elapsed = duration_since_epoch(last_modified);
410	let millis = elapsed
411		.as_secs()
412		.saturating_mul(MILLIS_PER_SECOND)
413		.saturating_add(u64::from(elapsed.subsec_millis()));
414
415	let timestamp = millis.to_be_bytes();
416	let digest = concat([id.as_bytes(), data.as_ref(), timestamp.as_slice()].into_iter());
417	let mut encoded = [0_u8; ETAG_VALUE_LENGTH];
418	let len = b64
419		.encode_slice(digest, &mut encoded)
420		.expect("ETag buffer has exact capacity");
421
422	let encoded = str::from_utf8(&encoded[..len]).expect("base64url is valid UTF-8");
423
424	Etag::try_from(format_args!("\"{encoded}\"")).expect("ETag has exact capacity")
425}
426
427fn get_outcome(session: &Session, if_none_match: Option<&str>) -> Get {
428	let meta = session.meta();
429
430	if if_none_match.is_some_and(|candidate| etag_matches(candidate, &session.etag)) {
431		Get::NotModified(meta)
432	} else {
433		Get::Data { data: session.data.clone(), meta }
434	}
435}
436
437fn etag_matches(candidate: &str, etag: &Etag) -> bool {
438	candidate == "*" || candidate == etag.as_str()
439}