1use std::time::{Duration, SystemTime};
9
10use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD as b64encode};
11use ruma::thirdparty::Medium;
12use serde::{Deserialize, Serialize};
13use subtle::ConstantTimeEq;
14use tuwunel_core::{
15 Err, Result, implement,
16 result::NotFound,
17 smallstr::SmallString,
18 utils::{
19 self,
20 hash::sha256,
21 time::{timepoint_from_now, timepoint_has_passed},
22 },
23};
24use tuwunel_database::{Cbor, Deserialized};
25
26use super::{Association, UiaaKey};
27
28type ClaimSid = SmallString<[u8; 43]>;
29
30const TOKEN_LENGTH: usize = 48;
32
33const MAX_VERIFY_ATTEMPTS: u32 = 5;
37
38const UIAA_SESSION_TTL: Duration = Duration::from_hours(24);
40
41#[derive(Clone, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
43enum PendingUse {
44 #[default]
45 Available,
46 Claimed(Box<UiaaKey>),
47 Spent,
48}
49
50#[derive(Clone, Debug, Deserialize, Serialize)]
57struct Pending {
58 client_secret: String,
59 medium: Medium,
60 address: String,
61 token: String,
62 send_attempt: u64,
63 attempts: u32,
64 validated_at: Option<SystemTime>,
65 expires_at: Option<SystemTime>,
66 #[serde(default)]
67 use_state: PendingUse,
68}
69
70#[derive(Clone, Debug)]
76pub struct PendingOutcome {
77 pub sid: String,
79
80 pub freshly_minted_token: Option<String>,
82}
83
84#[implement(super::Service)]
91#[tracing::instrument(level = "debug", skip(self, client_secret))]
92pub async fn create_or_reuse_pending(
93 &self,
94 client_secret: &str,
95 medium: Medium,
96 address: &str,
97 send_attempt: u64,
98 ttl: Duration,
99) -> Result<PendingOutcome> {
100 let sid = derive_sid(&medium, address, client_secret);
101 let _pending_lock = self.pending_mutex.lock(&sid).await;
102
103 match self.get_pending(&sid).await {
104 | Err(error) if error.is_not_found() => (),
105 | Err(error) => return Err(error),
106 | Ok(existing) if expired(&existing) => {
107 self.delete_pending_state(&sid, &existing).await?;
108 },
109 | Ok(existing) => {
110 if !matches!(existing.use_state, PendingUse::Available) {
111 return Err!(Request(ThreepidAuthFailed(
112 "The verification session has already been used"
113 )));
114 }
115
116 if existing.validated_at.is_none() && send_attempt <= existing.send_attempt {
117 return Ok(PendingOutcome { sid, freshly_minted_token: None });
118 }
119 },
120 }
121
122 let token = utils::random_string(TOKEN_LENGTH);
123 let expires_at = Some(timepoint_from_now(ttl)?);
124 let pending = Pending {
125 client_secret: client_secret.to_owned(),
126 medium,
127 address: address.to_owned(),
128 token: token.clone(),
129 send_attempt,
130 attempts: 0,
131 validated_at: None,
132 expires_at,
133 use_state: PendingUse::Available,
134 };
135
136 self.persist_pending(&sid, &pending);
137
138 Ok(PendingOutcome { sid, freshly_minted_token: Some(token) })
139}
140
141#[implement(super::Service)]
147#[tracing::instrument(level = "debug", skip(self, client_secret, token))]
148pub async fn validate_pending_token(
149 &self,
150 sid: &str,
151 client_secret: &str,
152 token: &str,
153) -> Result<()> {
154 let _pending_lock = self.pending_mutex.lock(sid).await;
155 let pending = self.get_pending(sid).await?;
156
157 if expired(&pending) {
158 self.delete_pending_state(sid, &pending).await?;
159
160 return Err!(Request(NotFound("The verification session has expired")));
161 }
162
163 if !matches!(pending.use_state, PendingUse::Available) {
164 return Err!(Request(ThreepidAuthFailed(
165 "The verification session has already been used"
166 )));
167 }
168
169 if pending.validated_at.is_some() {
170 return Err!(Request(ThreepidAuthFailed(
171 "The verification session has already been validated"
172 )));
173 }
174
175 let secret_ok = ct_eq(&pending.client_secret, client_secret);
176 let token_ok = ct_eq(&pending.token, token);
177
178 if !secret_ok || !token_ok {
179 let attempts = pending.attempts.saturating_add(1);
180 match attempts >= MAX_VERIFY_ATTEMPTS {
181 | true => self.delete_pending_state(sid, &pending).await?,
182 | false => self.persist_pending(sid, &Pending { attempts, ..pending }),
183 }
184
185 return Err!(Request(ThreepidAuthFailed("Invalid verification token")));
186 }
187
188 let validated_at = Some(SystemTime::now());
189 self.persist_pending(sid, &Pending { validated_at, ..pending });
190
191 Ok(())
192}
193
194#[implement(super::Service)]
201#[tracing::instrument(level = "debug", skip(self, client_secret, claim))]
202pub async fn claim_validated(
203 &self,
204 sid: &str,
205 client_secret: &str,
206 claim: UiaaKey,
207) -> Result<bool> {
208 let _pending_lock = self.pending_mutex.lock(sid).await;
209 let pending = match self.get_pending(sid).await {
210 | Ok(pending) => pending,
211 | Err(error) if error.is_not_found() => return Ok(false),
212 | Err(error) => return Err(error),
213 };
214
215 if expired(&pending) {
216 self.delete_pending_state(sid, &pending).await?;
217
218 return Ok(false);
219 }
220
221 if !ct_eq(&pending.client_secret, client_secret) {
222 return Ok(false);
223 }
224
225 if pending.validated_at.is_none() {
226 return Ok(false);
227 }
228
229 match &pending.use_state {
230 | PendingUse::Available => (),
231 | PendingUse::Claimed(owner) if owner.as_ref() == &claim => (),
232 | PendingUse::Claimed(_) | PendingUse::Spent => return Ok(false),
233 }
234
235 let _claim_lock = self.claim_mutex.lock(&claim).await;
236
237 if self
238 .claim_sid(&claim)
239 .await?
240 .is_some_and(|claimed_sid| claimed_sid != sid)
241 {
242 return Ok(false);
243 }
244
245 let expires_at = Some(timepoint_from_now(UIAA_SESSION_TTL)?).max(pending.expires_at);
246 let mut txn = self.db.database.txn();
247
248 txn.put_raw(&self.db.userdevicesessionid_threepid, &claim, sid);
249
250 let pending = Pending {
251 expires_at,
252 use_state: PendingUse::Claimed(Box::new(claim)),
253 ..pending
254 };
255
256 txn.raw_put(&self.db.threepidsid_pending, sid, Cbor(&pending));
257 txn.execute();
258
259 Ok(true)
260}
261
262#[implement(super::Service)]
269#[tracing::instrument(level = "debug", skip(self, claim))]
270pub async fn refresh_claim(&self, claim: &UiaaKey) -> Result<bool> {
271 let Some(sid) = self.claim_sid(claim).await? else {
272 return Ok(false);
273 };
274
275 let _pending_lock = self.pending_mutex.lock(sid.as_str()).await;
276 let _claim_lock = self.claim_mutex.lock(claim).await;
277
278 if self.claim_sid(claim).await?.as_deref() != Some(sid.as_str()) {
279 return Ok(false);
280 }
281
282 let pending = match self.get_pending(&sid).await {
283 | Ok(pending) => pending,
284 | Err(error) if error.is_not_found() => {
285 self.delete_claim_index(claim);
286
287 return Ok(false);
288 },
289 | Err(error) => return Err(error),
290 };
291
292 if expired(&pending) {
293 self.delete_pending_rows(&sid, Some(claim));
294
295 return Ok(false);
296 }
297
298 if !matches!(&pending.use_state, PendingUse::Claimed(owner) if owner.as_ref() == claim) {
299 self.delete_claim_index(claim);
300
301 return Ok(false);
302 }
303
304 let expires_at = Some(timepoint_from_now(UIAA_SESSION_TTL)?).max(pending.expires_at);
305 let pending = Pending { expires_at, ..pending };
306 let mut txn = self.db.database.txn();
307
308 txn.raw_put(&self.db.threepidsid_pending, &sid, Cbor(&pending));
309 txn.put_raw(&self.db.userdevicesessionid_threepid, claim, &sid);
310 txn.execute();
311
312 Ok(true)
313}
314
315#[implement(super::Service)]
321#[tracing::instrument(level = "debug", skip(self, claim))]
322pub async fn redeem_claim(&self, claim: &UiaaKey) -> Result<Association> {
323 let sid = self
324 .db
325 .userdevicesessionid_threepid
326 .qry(claim)
327 .await
328 .deserialized::<ClaimSid>()?;
329
330 let _pending_lock = self.pending_mutex.lock(sid.as_str()).await;
331 let _claim_lock = self.claim_mutex.lock(claim).await;
332 let current_sid = self
333 .db
334 .userdevicesessionid_threepid
335 .qry(claim)
336 .await
337 .deserialized::<ClaimSid>()?;
338
339 if current_sid != sid {
340 return Err!(Request(ThreepidAuthFailed("The verification session claim has changed")));
341 }
342
343 let pending = match self.get_pending(&sid).await {
344 | Ok(pending) => pending,
345 | Err(error) if error.is_not_found() => {
346 self.delete_claim_index(claim);
347
348 return Err(error);
349 },
350 | Err(error) => return Err(error),
351 };
352
353 if expired(&pending) {
354 self.delete_pending_rows(&sid, Some(claim));
355
356 return Err!(Request(NotFound("The verification session has expired")));
357 }
358
359 if !matches!(&pending.use_state, PendingUse::Claimed(owner) if owner.as_ref() == claim) {
360 self.delete_claim_index(claim);
361
362 return Err!(Request(ThreepidAuthFailed(
363 "The verification session is not owned by this transaction"
364 )));
365 }
366
367 let association = Association {
368 medium: pending.medium.clone(),
369 address: pending.address.clone(),
370 };
371
372 let pending = Pending { use_state: PendingUse::Spent, ..pending };
373 let mut txn = self.db.database.txn();
374
375 txn.raw_put(&self.db.threepidsid_pending, &sid, Cbor(&pending));
376 txn.del(&self.db.userdevicesessionid_threepid, claim);
377 txn.execute();
378
379 Ok(association)
380}
381
382#[implement(super::Service)]
390#[tracing::instrument(level = "debug", skip(self, client_secret))]
391pub async fn redeem_validated(&self, sid: &str, client_secret: &str) -> Result<Association> {
392 let _pending_lock = self.pending_mutex.lock(sid).await;
393 let pending = self.get_pending(sid).await?;
394
395 if expired(&pending) {
396 self.delete_pending_state(sid, &pending).await?;
397
398 return Err!(Request(NotFound("The verification session has expired")));
399 }
400
401 if !ct_eq(&pending.client_secret, client_secret) {
402 return Err!(Request(ThreepidAuthFailed("Client secret does not match")));
403 }
404
405 if pending.validated_at.is_none() {
406 return Err!(Request(ThreepidAuthFailed("The address has not been validated")));
407 }
408
409 if !matches!(pending.use_state, PendingUse::Available) {
410 return Err!(Request(ThreepidAuthFailed(
411 "The verification session has already been used"
412 )));
413 }
414
415 let association = Association {
416 medium: pending.medium.clone(),
417 address: pending.address.clone(),
418 };
419
420 self.persist_pending(sid, &Pending { use_state: PendingUse::Spent, ..pending });
421
422 Ok(association)
423}
424
425#[implement(super::Service)]
431#[tracing::instrument(level = "debug", skip(self, client_secret))]
432pub async fn session_validated(&self, sid: &str, client_secret: &str) -> bool {
433 let Ok(pending) = self.get_pending(sid).await else {
434 return false;
435 };
436
437 !expired(&pending)
438 && ct_eq(&pending.client_secret, client_secret)
439 && pending.validated_at.is_some()
440 && matches!(pending.use_state, PendingUse::Available)
441}
442
443#[implement(super::Service)]
444fn persist_pending(&self, sid: &str, pending: &Pending) {
445 self.db
446 .threepidsid_pending
447 .raw_put(sid, Cbor(pending));
448}
449
450#[implement(super::Service)]
451async fn delete_pending_state(&self, sid: &str, pending: &Pending) -> Result<()> {
452 let PendingUse::Claimed(claim) = &pending.use_state else {
453 self.delete_pending_rows(sid, None);
454
455 return Ok(());
456 };
457
458 let claim = claim.as_ref();
459 let _claim_lock = self.claim_mutex.lock(claim).await;
460 let claim = self
461 .claim_sid(claim)
462 .await?
463 .as_deref()
464 .is_some_and(|claimed_sid| claimed_sid == sid)
465 .then_some(claim);
466
467 self.delete_pending_rows(sid, claim);
468
469 Ok(())
470}
471
472#[implement(super::Service)]
473fn delete_pending_rows(&self, sid: &str, claim: Option<&UiaaKey>) {
474 let mut txn = self.db.database.txn();
475 txn.del_raw(&self.db.threepidsid_pending, sid);
476
477 if let Some(claim) = claim {
478 txn.del(&self.db.userdevicesessionid_threepid, claim);
479 }
480
481 txn.execute();
482}
483
484#[implement(super::Service)]
485fn delete_claim_index(&self, claim: &UiaaKey) { self.db.userdevicesessionid_threepid.del(claim); }
486
487#[implement(super::Service)]
488async fn claim_sid(&self, claim: &UiaaKey) -> Result<Option<ClaimSid>> {
489 self.db
490 .userdevicesessionid_threepid
491 .qry(claim)
492 .await
493 .deserialized::<ClaimSid>()
494 .optional()
495}
496
497#[implement(super::Service)]
498async fn get_pending(&self, sid: &str) -> Result<Pending> {
499 self.db
500 .threepidsid_pending
501 .get(sid)
502 .await
503 .deserialized::<Cbor<_>>()
504 .map(|Cbor(pending)| pending)
505}
506
507fn derive_sid(medium: &Medium, address: &str, client_secret: &str) -> String {
509 let parts = [medium.as_str().as_bytes(), address.as_bytes(), client_secret.as_bytes()];
510 let digest = sha256::delimited(parts.into_iter());
511
512 b64encode.encode(digest)
513}
514
515fn expired(pending: &Pending) -> bool {
516 pending
517 .expires_at
518 .is_some_and(timepoint_has_passed)
519}
520
521fn ct_eq(a: &str, b: &str) -> bool { a.as_bytes().ct_eq(b.as_bytes()).into() }
522
523#[cfg(test)]
524mod tests {
525 use std::time::SystemTime;
526
527 use ruma::{device_id, thirdparty::Medium, user_id};
528 use serde::Serialize;
529 use tuwunel_database::{Cbor, deserialize_from_slice, serialize_to_vec};
530
531 use super::{Pending, PendingUse};
532
533 #[derive(Serialize)]
534 struct LegacyPending {
535 client_secret: String,
536 medium: Medium,
537 address: String,
538 token: String,
539 send_attempt: u64,
540 attempts: u32,
541 validated_at: Option<SystemTime>,
542 expires_at: Option<SystemTime>,
543 }
544
545 fn pending(use_state: PendingUse) -> Pending {
546 Pending {
547 client_secret: "secret".into(),
548 medium: Medium::Email,
549 address: "user@example.com".into(),
550 token: "token".into(),
551 send_attempt: 1,
552 attempts: 0,
553 validated_at: None,
554 expires_at: None,
555 use_state,
556 }
557 }
558
559 fn round_trip(pending: Pending) -> Pending {
560 let encoded = serialize_to_vec(Cbor(pending)).expect("pending row should serialize");
561 let Cbor(pending): Cbor<Pending> =
562 deserialize_from_slice(&encoded).expect("pending row should deserialize");
563
564 pending
565 }
566
567 #[test]
568 fn legacy_pending_defaults_to_available() {
569 let legacy = LegacyPending {
570 client_secret: "secret".into(),
571 medium: Medium::Email,
572 address: "user@example.com".into(),
573 token: "token".into(),
574 send_attempt: 1,
575 attempts: 0,
576 validated_at: None,
577 expires_at: None,
578 };
579
580 let encoded =
581 serialize_to_vec(Cbor(legacy)).expect("legacy pending row should serialize");
582
583 let Cbor(pending): Cbor<Pending> =
584 deserialize_from_slice(&encoded).expect("legacy pending row should deserialize");
585
586 assert_eq!(pending.use_state, PendingUse::Available);
587 }
588
589 #[test]
590 fn claimed_pending_round_trip_preserves_exact_key() {
591 let claim = (
592 user_id!("@owner:example.org").to_owned(),
593 device_id!("DEVICE").to_owned(),
594 "0123456789abcdefghijklmnopqrstuv".into(),
595 );
596 let pending = round_trip(pending(PendingUse::Claimed(Box::new(claim.clone()))));
597
598 assert_eq!(pending.use_state, PendingUse::Claimed(Box::new(claim)));
599 }
600
601 #[test]
602 fn spent_pending_round_trip_preserves_tombstone() {
603 let pending = round_trip(pending(PendingUse::Spent));
604
605 assert_eq!(pending.use_state, PendingUse::Spent);
606 }
607}