1use std::{
7 cmp::max,
8 collections::{BTreeMap, HashMap},
9 net::IpAddr,
10 str,
11 sync::{Arc, Mutex, RwLock},
12 time::{Duration, Instant, SystemTime, UNIX_EPOCH},
13};
14
15use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD as b64};
16use bytes::Bytes;
17use http::StatusCode;
18use ruma::api::error::{ErrorKind, LimitExceededErrorData};
19use tuwunel_core::{
20 Error, Result,
21 arrayvec::ArrayString,
22 implement,
23 utils::{hash::sha256::concat, rand::string_array, time::duration_since_epoch},
24};
25
26pub type SessionId = ArrayString<SESSION_ID_LENGTH>;
30
31pub type Etag = ArrayString<ETAG_LENGTH>;
35type Sessions = BTreeMap<SessionId, Session>;
36type Ratelimiter = Mutex<HashMap<IpAddr, (Instant, f64)>>;
37
38pub struct Service {
43 sessions: RwLock<Sessions>,
44 ratelimiter: Ratelimiter,
46 services: Arc<crate::services::OnceServices>,
47}
48
49struct Session {
54 data: Bytes,
55 etag: Etag,
56 created: SystemTime,
57 last_modified: SystemTime,
58 expires_at: SystemTime,
59}
60
61#[derive(Clone, Copy, Debug, Eq, PartialEq)]
66pub struct Meta {
67 pub etag: Etag,
69
70 pub expires_at: SystemTime,
72
73 pub last_modified: SystemTime,
75}
76
77#[derive(Clone, Debug, Eq, PartialEq)]
82pub enum Get {
83 Data {
85 data: Bytes,
87
88 meta: Meta,
90 },
91
92 NotModified(Meta),
94
95 NotFound,
97}
98
99#[derive(Clone, Copy, Debug, Eq, PartialEq)]
104pub enum Put {
105 Accepted(Meta),
107
108 PreconditionFailed(Meta),
110
111 NotFound,
113}
114
115#[derive(Clone, Copy)]
120enum Validator<'a> {
121 Etag(&'a str),
123
124 SequenceToken(&'a str),
126}
127
128const SESSION_ID_LENGTH: usize = 32;
129const ETAG_VALUE_LENGTH: usize = 43;
130const ETAG_LENGTH: usize = ETAG_VALUE_LENGTH + 2;
131const MILLIS_PER_SECOND: u64 = 1000;
132const MAX_HTTP_DATE_SECONDS: u64 = 253_402_300_799;
133const MONOTONIC_STEP: Duration = Duration::from_millis(1);
134const RATELIMIT_MAP_CAP: usize = 4096;
135
136#[cfg(test)]
137mod tests;
138
139impl crate::Service for Service {
140 fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
141 Ok(Arc::new(Self {
142 sessions: RwLock::new(Sessions::new()),
143 ratelimiter: Mutex::new(HashMap::new()),
144 services: args.services.clone(),
145 }))
146 }
147
148 fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
149}
150
151#[implement(Service)]
156pub fn check_rate_limit(&self, client: IpAddr) -> Result {
157 let config = &self.services.server.config;
158 let rate = f64::from(config.rendezvous_rc_per_second.max(1));
159 let burst = f64::from(config.rendezvous_rc_burst_count.max(1));
160
161 check_bucket_at(&self.ratelimiter, client, rate, burst, Instant::now())
162}
163
164fn check_bucket_at(
165 table: &Ratelimiter,
166 client: IpAddr,
167 rate: f64,
168 burst: f64,
169 now: Instant,
170) -> Result {
171 let mut buckets = table.lock()?;
172
173 if buckets.len() >= RATELIMIT_MAP_CAP && !buckets.contains_key(&client) {
174 let mut oldest = None;
175
176 buckets.retain(|client, (last, tokens)| {
177 let refilled = now
178 .duration_since(*last)
179 .as_secs_f64()
180 .mul_add(rate, *tokens);
181 let retain = refilled < burst;
182
183 if retain && oldest.is_none_or(|(_, oldest_at)| *last < oldest_at) {
184 oldest = Some((*client, *last));
185 }
186
187 retain
188 });
189
190 if buckets.len() >= RATELIMIT_MAP_CAP
191 && let Some((oldest, _)) = oldest
192 {
193 buckets.remove(&oldest);
194 }
195 }
196
197 let (last_time, tokens) = buckets.entry(client).or_insert((now, burst));
198 let new_tokens = now
199 .duration_since(*last_time)
200 .as_secs_f64()
201 .mul_add(rate, *tokens)
202 .min(burst);
203
204 if new_tokens < 1.0 {
205 return Err(Error::Request(
206 ErrorKind::LimitExceeded(LimitExceededErrorData { retry_after: None }),
207 "Too many rendezvous requests.".into(),
208 StatusCode::TOO_MANY_REQUESTS,
209 ));
210 }
211
212 *last_time = now;
213 *tokens = new_tokens - 1.0;
214
215 Ok(())
216}
217
218#[implement(Service)]
228pub fn create(&self, data: Bytes) -> (SessionId, Meta) {
229 let config = &self.services.server.config;
230 let ttl = Duration::from_secs(config.rendezvous_session_ttl);
231
232 self.create_at(data, SystemTime::now(), ttl, config.rendezvous_max_sessions)
233}
234
235#[implement(Service)]
236fn create_at(
237 &self,
238 data: Bytes,
239 now: SystemTime,
240 ttl: Duration,
241 max_sessions: usize,
242) -> (SessionId, Meta) {
243 let capacity = max_sessions.max(1);
244 let mut id = string_array::<SESSION_ID_LENGTH>();
245 let expires_at = expires_at(now, ttl);
246
247 let mut sessions = self.sessions.write().expect("locked for writing");
248
249 sessions.retain(|_, session| session.expires_at > now);
250 while sessions.contains_key(id.as_str()) {
251 id = string_array::<SESSION_ID_LENGTH>();
252 }
253
254 while sessions.len() >= capacity {
255 let Some(id) = sessions
256 .iter()
257 .min_by_key(|(_, session)| session.created)
258 .map(|(id, _)| *id)
259 else {
260 break;
261 };
262
263 sessions.remove(&id);
264 }
265
266 let session = Session {
267 etag: etag(id.as_str(), &data, now),
268 data,
269 created: now,
270 last_modified: now,
271 expires_at,
272 };
273
274 let meta = session.meta();
275
276 sessions.insert(id, session);
277
278 (id, meta)
279}
280
281#[implement(Service)]
290pub fn get(&self, id: &str, if_none_match: Option<&str>) -> Get {
291 self.get_at(id, if_none_match, SystemTime::now())
292}
293
294#[implement(Service)]
295fn get_at(&self, id: &str, if_none_match: Option<&str>, now: SystemTime) -> Get {
296 {
297 let sessions = self.sessions.read().expect("locked for reading");
298
299 match sessions.get(id) {
300 | Some(session) if session.expires_at > now => {
301 return get_outcome(session, if_none_match);
302 },
303 | Some(_) => {},
304 | None => return Get::NotFound,
305 }
306 }
307
308 let mut sessions = self.sessions.write().expect("locked for writing");
309 if sessions
310 .get(id)
311 .is_some_and(|session| session.expires_at <= now)
312 {
313 sessions.remove(id);
314
315 return Get::NotFound;
316 }
317
318 sessions
319 .get(id)
320 .map_or(Get::NotFound, |session| get_outcome(session, if_none_match))
321}
322
323#[implement(Service)]
332pub fn put(&self, id: &str, if_match: &str, data: Bytes) -> Put {
333 let ttl = Duration::from_secs(self.services.server.config.rendezvous_session_ttl);
334
335 self.put_at(id, if_match, data, SystemTime::now(), ttl)
336}
337
338#[implement(Service)]
339fn put_at(&self, id: &str, if_match: &str, data: Bytes, now: SystemTime, ttl: Duration) -> Put {
340 self.put_with_at(id, Validator::Etag(if_match), data, now, ttl)
341}
342
343#[implement(Service)]
352pub fn put_token(&self, id: &str, sequence_token: &str, data: Bytes) -> Put {
353 let ttl = Duration::from_secs(self.services.server.config.rendezvous_session_ttl);
354
355 self.put_token_at(id, sequence_token, data, SystemTime::now(), ttl)
356}
357
358#[implement(Service)]
359fn put_token_at(
360 &self,
361 id: &str,
362 sequence_token: &str,
363 data: Bytes,
364 now: SystemTime,
365 ttl: Duration,
366) -> Put {
367 self.put_with_at(id, Validator::SequenceToken(sequence_token), data, now, ttl)
368}
369
370#[implement(Service)]
375fn put_with_at(
376 &self,
377 id: &str,
378 validator: Validator<'_>,
379 data: Bytes,
380 now: SystemTime,
381 ttl: Duration,
382) -> Put {
383 let mut sessions = self.sessions.write().expect("locked for writing");
384
385 if sessions
386 .get(id)
387 .is_some_and(|session| session.expires_at <= now)
388 {
389 sessions.remove(id);
390
391 return Put::NotFound;
392 }
393
394 let Some(session) = sessions.get_mut(id) else {
395 return Put::NotFound;
396 };
397
398 if !validator.matches(&session.etag) {
399 if data != session.data {
400 return Put::PreconditionFailed(session.meta());
401 }
402
403 session.expires_at = expires_at(now, ttl);
404
405 return Put::Accepted(session.meta());
406 }
407
408 let created = session.created;
409 let last_modified = next_last_modified(now, session.last_modified);
410 let expires_at = expires_at(now, ttl);
411
412 *session = Session {
413 etag: etag(id, &data, last_modified),
414 data,
415 created,
416 last_modified,
417 expires_at,
418 };
419
420 Put::Accepted(session.meta())
421}
422
423#[implement(Service)]
431pub fn delete(&self, id: &str) -> bool {
432 self.sessions
433 .write()
434 .expect("locked for writing")
435 .remove(id)
436 .is_some()
437}
438
439#[implement(Service)]
447pub fn delete_if_active(&self, id: &str) -> bool {
448 self.delete_if_active_at(id, SystemTime::now())
449}
450
451#[implement(Service)]
452fn delete_if_active_at(&self, id: &str, now: SystemTime) -> bool {
453 self.sessions
454 .write()
455 .expect("locked for writing")
456 .remove(id)
457 .is_some_and(|session| session.expires_at > now)
458}
459
460#[implement(Meta)]
468#[must_use]
469#[inline]
470pub fn sequence_token(&self) -> &str { etag_value(&self.etag) }
471
472fn etag_value(etag: &Etag) -> &str {
473 etag.as_str()
474 .strip_prefix('"')
475 .and_then(|value| value.strip_suffix('"'))
476 .expect("ETag is quoted")
477}
478
479#[implement(Meta)]
483#[must_use]
484#[inline]
485pub fn expires_in(&self) -> Duration { self.expires_in_at(SystemTime::now()) }
486
487#[implement(Meta)]
488fn expires_in_at(&self, now: SystemTime) -> Duration {
489 self.expires_at
490 .duration_since(now)
491 .unwrap_or_default()
492}
493
494impl Session {
495 fn meta(&self) -> Meta {
496 Meta {
497 etag: self.etag,
498 expires_at: self.expires_at,
499 last_modified: self.last_modified,
500 }
501 }
502}
503
504impl Validator<'_> {
505 fn matches(self, stored: &Etag) -> bool {
506 match self {
507 | Self::Etag(candidate) => etag_matches(candidate, stored),
508 | Self::SequenceToken(candidate) => candidate == etag_value(stored),
509 }
510 }
511}
512
513fn expires_at(now: SystemTime, ttl: Duration) -> SystemTime {
514 let latest = UNIX_EPOCH
515 .checked_add(Duration::from_secs(MAX_HTTP_DATE_SECONDS))
516 .expect("latest HTTP date should fit in SystemTime");
517
518 now.checked_add(ttl)
519 .map_or(latest, |expires_at| expires_at.min(latest))
520}
521
522fn next_last_modified(now: SystemTime, previous: SystemTime) -> SystemTime {
523 previous
524 .checked_add(MONOTONIC_STEP)
525 .map_or(now, |next| max(now, next))
526}
527
528fn etag(id: &str, data: &Bytes, last_modified: SystemTime) -> Etag {
529 let elapsed = duration_since_epoch(last_modified);
530 let millis = elapsed
531 .as_secs()
532 .saturating_mul(MILLIS_PER_SECOND)
533 .saturating_add(u64::from(elapsed.subsec_millis()));
534
535 let timestamp = millis.to_be_bytes();
536 let digest = concat([id.as_bytes(), data.as_ref(), timestamp.as_slice()].into_iter());
537 let mut encoded = [0_u8; ETAG_VALUE_LENGTH];
538 let len = b64
539 .encode_slice(digest, &mut encoded)
540 .expect("ETag buffer has exact capacity");
541
542 let encoded = str::from_utf8(&encoded[..len]).expect("base64url is valid UTF-8");
543
544 Etag::try_from(format_args!("\"{encoded}\"")).expect("ETag has exact capacity")
545}
546
547fn get_outcome(session: &Session, if_none_match: Option<&str>) -> Get {
548 let meta = session.meta();
549
550 if if_none_match.is_some_and(|candidate| etag_matches(candidate, &session.etag)) {
551 Get::NotModified(meta)
552 } else {
553 Get::Data { data: session.data.clone(), meta }
554 }
555}
556
557fn etag_matches(candidate: &str, etag: &Etag) -> bool {
558 candidate == "*" || candidate == etag.as_str()
559}