1use 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 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}